mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 09:57:00 +02:00
Compare commits
105 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 7fe4fab167 | |||
| 37b924945d | |||
| 5631346b07 | |||
| 97eb3c635a | |||
| 05f7923860 | |||
| 49028cb755 | |||
| adf71d1458 | |||
| cefda6524c | |||
| d0f14de739 | |||
| 9e7646ee62 | |||
| 6e6cfc89db | |||
| 3719f135e3 | |||
| 2724b4a96c | |||
| e386e290ab | |||
| 0a44376403 | |||
| 9e61269935 | |||
| a0836aa819 | |||
| cf2800f7bd | |||
| fc89b9d14c | |||
| 3e5326e60a | |||
| 840841a53c | |||
| cfb4ab24c3 | |||
| 67f9cfad91 | |||
| 1c78f5e500 | |||
| ad8ff2e45e | |||
| c13a1ff4ec | |||
| 8b66ebaf72 | |||
| bbdaf5c3f3 | |||
| a515aeaf39 | |||
| 0f6f14eaf6 | |||
| dafdd34af9 | |||
| f6791130df | |||
| 145a6267fa | |||
| 7716f1da23 | |||
| beb1d7d89f | |||
| fa0d593f28 | |||
| ff48040c78 | |||
| 6c3972f464 | |||
| 861d3aabd7 | |||
| 86733864fe | |||
| ab736e4c6b | |||
| 5ecdd4eaa9 | |||
| 1b84bd0050 | |||
| 384610f81a | |||
| c1eea118f4 | |||
| 1e66c0d3ee | |||
| e5c0fdad8d | |||
| 7bd0bc285a | |||
| 942ee522e0 | |||
| 32149f3a93 | |||
| 19ad3bb904 | |||
| 647775d8c3 | |||
| abfeb502a8 | |||
| 0a953915bb | |||
| 6aa3363d85 | |||
| 95d98b1f4b | |||
| 6afca0f461 | |||
| 02471b4121 | |||
| 58ab7250f5 | |||
| 184cdc8586 | |||
| 7d3166a19d | |||
| fe1c5682f0 | |||
| e4cc80aaca | |||
| 16b302c11d | |||
| ab539f8a3f | |||
| b7d83b0500 | |||
| ef95b25fa3 | |||
| 36b38396af | |||
| 2e9117da5b | |||
| a690c904ba | |||
| e028e6bf1a | |||
| 3db406b8ac | |||
| eaad4896c1 | |||
| e6032f81aa | |||
| b041f306cb | |||
| 3a95495c63 | |||
| 873f94f465 | |||
| 72bad1603a | |||
| 0c1ad9bb48 | |||
| 074a123a4b | |||
| 04dea41f74 | |||
| 0d23377c65 | |||
| ffd5249cf5 | |||
| 625f58b84a | |||
| 99c5854e5c | |||
| 3c121e7ab1 | |||
| 6c7ebb0875 | |||
| 110ea8f45c | |||
| 398d67e2da | |||
| 696903d6d9 | |||
| c82db210ef | |||
| 1ada3d4dd9 | |||
| 5f920fdd7d | |||
| cba9ea5b1f | |||
| 83809a599a | |||
| 23c67bd8d8 | |||
| dd3a7ad03c | |||
| dd2ac5d655 | |||
| 76e82a5256 | |||
| eaf756ea6c | |||
| a82a8dc547 | |||
| 213dd46588 | |||
| 4fb5cdb4fa | |||
| ff91c37529 | |||
| b7e9939e92 |
@@ -0,0 +1,116 @@
|
|||||||
|
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 to write under; defaults to code-signing/<owner>/<repo> of the calling repo"
|
||||||
|
required: false
|
||||||
|
default: ""
|
||||||
|
|
||||||
|
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
|
||||||
|
# Default the prefix to this repo so the S3 key attributes the sign correctly.
|
||||||
|
# nebula-nightly runs this same action but writes under its own repo's prefix.
|
||||||
|
KEY_PREFIX="${KEY_PREFIX:-code-signing/$GITHUB_REPOSITORY}"
|
||||||
|
RUN="${GITHUB_RUN_ID}-${GITHUB_RUN_ATTEMPT}"
|
||||||
|
|
||||||
|
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
|
||||||
@@ -1,34 +0,0 @@
|
|||||||
name: gofmt
|
|
||||||
on:
|
|
||||||
push:
|
|
||||||
branches:
|
|
||||||
- master
|
|
||||||
pull_request:
|
|
||||||
paths:
|
|
||||||
- '.github/workflows/gofmt.yml'
|
|
||||||
- '**.go'
|
|
||||||
jobs:
|
|
||||||
|
|
||||||
gofmt:
|
|
||||||
name: Run gofmt
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
steps:
|
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
|
||||||
with:
|
|
||||||
go-version: '1.25'
|
|
||||||
check-latest: true
|
|
||||||
|
|
||||||
- name: Install goimports
|
|
||||||
run: |
|
|
||||||
go install golang.org/x/tools/cmd/goimports@latest
|
|
||||||
|
|
||||||
- name: gofmt
|
|
||||||
run: |
|
|
||||||
if [ "$(find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -l)" ]
|
|
||||||
then
|
|
||||||
find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -d
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
@@ -10,7 +10,7 @@ jobs:
|
|||||||
name: Build Linux/BSD All
|
name: Build Linux/BSD All
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -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,8 +32,11 @@ 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@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -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
|
||||||
@@ -66,7 +76,7 @@ jobs:
|
|||||||
HAS_SIGNING_CREDS: ${{ secrets.AC_USERNAME != '' }}
|
HAS_SIGNING_CREDS: ${{ secrets.AC_USERNAME != '' }}
|
||||||
runs-on: macos-latest
|
runs-on: macos-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -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/*
|
||||||
@@ -124,25 +134,25 @@ jobs:
|
|||||||
# be overwritten
|
# be overwritten
|
||||||
- name: Checkout code
|
- name: Checkout code
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
uses: actions/checkout@v6
|
uses: actions/checkout@v7
|
||||||
|
|
||||||
- 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' }}
|
||||||
@@ -153,17 +163,20 @@ jobs:
|
|||||||
mkdir -p build/linux-{amd64,arm64}
|
mkdir -p build/linux-{amd64,arm64}
|
||||||
tar -zxvf artifacts/nebula-linux-amd64.tar.gz -C build/linux-amd64/
|
tar -zxvf artifacts/nebula-linux-amd64.tar.gz -C build/linux-amd64/
|
||||||
tar -zxvf artifacts/nebula-linux-arm64.tar.gz -C build/linux-arm64/
|
tar -zxvf artifacts/nebula-linux-arm64.tar.gz -C build/linux-arm64/
|
||||||
docker buildx build . --push -f docker/Dockerfile --platform linux/amd64,linux/arm64 --tag "${DOCKER_IMAGE_REPO}:${DOCKER_IMAGE_TAG}" --tag "${DOCKER_IMAGE_REPO}:${GITHUB_REF#refs/tags/v}"
|
docker buildx build . --push -f docker/Dockerfile --platform linux/amd64,linux/arm64 \
|
||||||
|
--build-arg VERSION="${GITHUB_REF#refs/tags/v}" \
|
||||||
|
--build-arg REVISION="${GITHUB_SHA}" \
|
||||||
|
--tag "${DOCKER_IMAGE_REPO}:${DOCKER_IMAGE_TAG}" --tag "${DOCKER_IMAGE_REPO}:${GITHUB_REF#refs/tags/v}"
|
||||||
|
|
||||||
release:
|
release:
|
||||||
name: Create and Upload Release
|
name: Create and Upload Release
|
||||||
needs: [build-linux, build-darwin, build-windows]
|
needs: [build-linux, build-darwin, build-windows]
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
uses: actions/download-artifact@v7
|
uses: actions/download-artifact@v8
|
||||||
with:
|
with:
|
||||||
path: artifacts
|
path: artifacts
|
||||||
|
|
||||||
|
|||||||
@@ -14,15 +14,23 @@ 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:
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -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@v7
|
||||||
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@v7
|
||||||
|
|
||||||
|
- 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
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -36,6 +36,14 @@ jobs:
|
|||||||
working-directory: ./.github/workflows/smoke
|
working-directory: ./.github/workflows/smoke
|
||||||
run: ./smoke.sh
|
run: ./smoke.sh
|
||||||
|
|
||||||
|
- name: setup docker image ipv6
|
||||||
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: SMOKE_OVERLAY_IPV6=1 ./build.sh
|
||||||
|
|
||||||
|
- name: run smoke ipv6
|
||||||
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: SMOKE_OVERLAY_IPV6=1 ./smoke.sh
|
||||||
|
|
||||||
- name: setup relay docker image
|
- name: setup relay docker image
|
||||||
working-directory: ./.github/workflows/smoke
|
working-directory: ./.github/workflows/smoke
|
||||||
run: ./build-relay.sh
|
run: ./build-relay.sh
|
||||||
|
|||||||
@@ -5,6 +5,19 @@ set -e -x
|
|||||||
rm -rf ./build
|
rm -rf ./build
|
||||||
mkdir ./build
|
mkdir ./build
|
||||||
|
|
||||||
|
if [ "$SMOKE_OVERLAY_IPV6" ]
|
||||||
|
then
|
||||||
|
LIGHTHOUSE_NIP="fd00:4242:0:0:0:ffff:c0a8:6401"
|
||||||
|
HOST2_NIP="fd00:4242:0:0:0:ffff:c0a8:6402"
|
||||||
|
HOST3_NIP="fd00:4242:0:0:0:ffff:c0a8:6403"
|
||||||
|
HOST4_NIP="fd00:4242:0:0:0:ffff:c0a8:6404"
|
||||||
|
else
|
||||||
|
LIGHTHOUSE_NIP="192.168.100.1"
|
||||||
|
HOST2_NIP="192.168.100.2"
|
||||||
|
HOST3_NIP="192.168.100.3"
|
||||||
|
HOST4_NIP="192.168.100.4"
|
||||||
|
fi
|
||||||
|
|
||||||
# Smoke containers run on a dedicated docker network whose subnet is allocated
|
# Smoke containers run on a dedicated docker network whose subnet is allocated
|
||||||
# at smoke time, not known at build time. Configs are written with TEST-NET-3
|
# at smoke time, not known at build time. Configs are written with TEST-NET-3
|
||||||
# placeholder IPs (RFC 5737) and smoke.sh / smoke-vagrant.sh / smoke-relay.sh
|
# placeholder IPs (RFC 5737) and smoke.sh / smoke-vagrant.sh / smoke-relay.sh
|
||||||
@@ -31,24 +44,24 @@ LIGHTHOUSE_IP="203.0.113.2"
|
|||||||
../genconfig.sh >lighthouse1.yml
|
../genconfig.sh >lighthouse1.yml
|
||||||
|
|
||||||
HOST="host2" \
|
HOST="host2" \
|
||||||
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \
|
||||||
../genconfig.sh >host2.yml
|
../genconfig.sh >host2.yml
|
||||||
|
|
||||||
HOST="host3" \
|
HOST="host3" \
|
||||||
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \
|
||||||
INBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
INBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
||||||
../genconfig.sh >host3.yml
|
../genconfig.sh >host3.yml
|
||||||
|
|
||||||
HOST="host4" \
|
HOST="host4" \
|
||||||
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \
|
||||||
OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
||||||
../genconfig.sh >host4.yml
|
../genconfig.sh >host4.yml
|
||||||
|
|
||||||
../../../../nebula-cert ca -curve "${CURVE:-25519}" -name "Smoke Test"
|
../../../../nebula-cert ca -curve "${CURVE:-25519}" -name "Smoke Test"
|
||||||
../../../../nebula-cert sign -name "lighthouse1" -groups "lighthouse,lighthouse1" -ip "192.168.100.1/24"
|
../../../../nebula-cert sign -name "lighthouse1" -groups "lighthouse,lighthouse1" -ip "$LIGHTHOUSE_NIP/24"
|
||||||
../../../../nebula-cert sign -name "host2" -groups "host,host2" -ip "192.168.100.2/24"
|
../../../../nebula-cert sign -name "host2" -groups "host,host2" -ip "$HOST2_NIP/24"
|
||||||
../../../../nebula-cert sign -name "host3" -groups "host,host3" -ip "192.168.100.3/24"
|
../../../../nebula-cert sign -name "host3" -groups "host,host3" -ip "$HOST3_NIP/24"
|
||||||
../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "192.168.100.4/24"
|
../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "$HOST4_NIP/24"
|
||||||
)
|
)
|
||||||
|
|
||||||
docker build -t "nebula:${NAME:-smoke}" .
|
docker build -t "nebula:${NAME:-smoke}" .
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -47,6 +47,19 @@ HOST2_IP="$PREFIX.3"
|
|||||||
HOST3_IP="$PREFIX.4"
|
HOST3_IP="$PREFIX.4"
|
||||||
HOST4_IP="$PREFIX.5"
|
HOST4_IP="$PREFIX.5"
|
||||||
|
|
||||||
|
if [ "$SMOKE_OVERLAY_IPV6" ]
|
||||||
|
then
|
||||||
|
LIGHTHOUSE_NIP="fd00:4242:0:0:0:ffff:c0a8:6401"
|
||||||
|
HOST2_NIP="fd00:4242:0:0:0:ffff:c0a8:6402"
|
||||||
|
HOST3_NIP="fd00:4242:0:0:0:ffff:c0a8:6403"
|
||||||
|
HOST4_NIP="fd00:4242:0:0:0:ffff:c0a8:6404"
|
||||||
|
else
|
||||||
|
LIGHTHOUSE_NIP="192.168.100.1"
|
||||||
|
HOST2_NIP="192.168.100.2"
|
||||||
|
HOST3_NIP="192.168.100.3"
|
||||||
|
HOST4_NIP="192.168.100.4"
|
||||||
|
fi
|
||||||
|
|
||||||
# Sed the placeholder TEST-NET-3 IPs in the host configs to the real ones.
|
# Sed the placeholder TEST-NET-3 IPs in the host configs to the real ones.
|
||||||
# build/lighthouse1.yml has no IPs to rewrite so it's skipped.
|
# build/lighthouse1.yml has no IPs to rewrite so it's skipped.
|
||||||
for f in build/host2.yml build/host3.yml build/host4.yml; do
|
for f in build/host2.yml build/host3.yml build/host4.yml; do
|
||||||
@@ -80,28 +93,28 @@ docker exec host3 tcpdump -i eth0 -q -w - -U 2>logs/host3.outside.log >logs/host
|
|||||||
docker exec host4 tcpdump -i tun0 -q -w - -U 2>logs/host4.inside.log >logs/host4.inside.pcap &
|
docker exec host4 tcpdump -i tun0 -q -w - -U 2>logs/host4.inside.log >logs/host4.inside.pcap &
|
||||||
docker exec host4 tcpdump -i eth0 -q -w - -U 2>logs/host4.outside.log >logs/host4.outside.pcap &
|
docker exec host4 tcpdump -i eth0 -q -w - -U 2>logs/host4.outside.log >logs/host4.outside.pcap &
|
||||||
|
|
||||||
docker exec host2 ncat -nklv 0.0.0.0 2000 &
|
docker exec host2 ncat -nklv 2000 &
|
||||||
docker exec host3 ncat -nklv 0.0.0.0 2000 &
|
docker exec host3 ncat -nklv 2000 &
|
||||||
docker exec host4 ncat -e '/usr/bin/echo helloagainfromhost4' -nkluv 0.0.0.0 4000 &
|
docker exec host4 ncat -e '/usr/bin/echo helloagainfromhost4' -nkluv 4000 &
|
||||||
docker exec host2 ncat -e '/usr/bin/echo host2' -nkluv 0.0.0.0 3000 &
|
docker exec host2 ncat -e '/usr/bin/echo host2' -nkluv 3000 &
|
||||||
docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 0.0.0.0 3000 &
|
docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 3000 &
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from lighthouse1"
|
echo " *** Testing ping from lighthouse1"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec lighthouse1 ping -c1 192.168.100.2
|
docker exec lighthouse1 ping -c1 $HOST2_NIP
|
||||||
docker exec lighthouse1 ping -c1 192.168.100.3
|
docker exec lighthouse1 ping -c1 $HOST3_NIP
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host2"
|
echo " *** Testing ping from host2"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host2 ping -c1 192.168.100.1
|
docker exec host2 ping -c1 $LIGHTHOUSE_NIP
|
||||||
# Should fail because not allowed by host3 inbound firewall
|
# Should fail because not allowed by host3 inbound firewall
|
||||||
! docker exec host2 ping -c1 192.168.100.3 -w5 || exit 1
|
! docker exec host2 ping -c1 $HOST3_NIP -w5 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
@@ -109,34 +122,34 @@ echo " *** Testing ncat from host2"
|
|||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
# Should fail because not allowed by host3 inbound firewall
|
# Should fail because not allowed by host3 inbound firewall
|
||||||
! docker exec host2 ncat -nzv -w5 192.168.100.3 2000 || exit 1
|
! docker exec host2 ncat -nzv -w5 $HOST3_NIP 2000 || exit 1
|
||||||
! docker exec host2 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
|
! docker exec host2 ncat -nzuv -w5 $HOST3_NIP 3000 | grep -q host3 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host3"
|
echo " *** Testing ping from host3"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host3 ping -c1 192.168.100.1
|
docker exec host3 ping -c1 $LIGHTHOUSE_NIP
|
||||||
docker exec host3 ping -c1 192.168.100.2
|
docker exec host3 ping -c1 $HOST2_NIP
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ncat from host3"
|
echo " *** Testing ncat from host3"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host3 ncat -nzv -w5 192.168.100.2 2000
|
docker exec host3 ncat -nzv -w5 $HOST2_NIP 2000
|
||||||
docker exec host3 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2
|
docker exec host3 ncat -nzuv -w5 $HOST2_NIP 3000 | grep -q host2
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host4"
|
echo " *** Testing ping from host4"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host4 ping -c1 192.168.100.1
|
docker exec host4 ping -c1 $LIGHTHOUSE_NIP
|
||||||
# Should fail because not allowed by host4 outbound firewall
|
# Should fail because not allowed by host4 outbound firewall
|
||||||
! docker exec host4 ping -c1 192.168.100.2 -w5 || exit 1
|
! docker exec host4 ping -c1 $HOST2_NIP -w5 || exit 1
|
||||||
! docker exec host4 ping -c1 192.168.100.3 -w5 || exit 1
|
! docker exec host4 ping -c1 $HOST3_NIP -w5 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
@@ -144,10 +157,10 @@ echo " *** Testing ncat from host4"
|
|||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
# Should fail because not allowed by host4 outbound firewall
|
# Should fail because not allowed by host4 outbound firewall
|
||||||
! docker exec host4 ncat -nzv -w5 192.168.100.2 2000 || exit 1
|
! docker exec host4 ncat -nzv -w5 $HOST2_NIP 2000 || exit 1
|
||||||
! docker exec host4 ncat -nzv -w5 192.168.100.3 2000 || exit 1
|
! docker exec host4 ncat -nzv -w5 $HOST3_NIP 2000 || exit 1
|
||||||
! docker exec host4 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2 || exit 1
|
! docker exec host4 ncat -nzuv -w5 $HOST2_NIP 3000 | grep -q host2 || exit 1
|
||||||
! docker exec host4 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
|
! docker exec host4 ncat -nzuv -w5 $HOST3_NIP 3000 | grep -q host3 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
@@ -159,7 +172,7 @@ set -x
|
|||||||
# cannot initiate UDP to host2. Once host2 initiates a flow to host4:4000,
|
# cannot initiate UDP to host2. Once host2 initiates a flow to host4:4000,
|
||||||
# conntrack must let host4's listener reply on that flow. If it doesn't,
|
# conntrack must let host4's listener reply on that flow. If it doesn't,
|
||||||
# the echo back from host4 never reaches host2.
|
# the echo back from host4 never reaches host2.
|
||||||
docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv 192.168.100.4 4000" | grep -q helloagainfromhost4
|
docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv $HOST4_NIP 4000" | grep -q helloagainfromhost4
|
||||||
|
|
||||||
docker exec host4 sh -c 'kill 1'
|
docker exec host4 sh -c 'kill 1'
|
||||||
docker exec host3 sh -c 'kill 1'
|
docker exec host3 sh -c 'kill 1'
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+97
-79
@@ -13,20 +13,28 @@ on:
|
|||||||
- 'go.sum'
|
- 'go.sum'
|
||||||
jobs:
|
jobs:
|
||||||
|
|
||||||
test-linux:
|
static:
|
||||||
name: Build all and test on ubuntu-linux
|
name: Static checks
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.25'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Install goimports
|
||||||
run: make all
|
run: go install golang.org/x/tools/cmd/goimports@latest
|
||||||
|
|
||||||
|
- name: gofmt
|
||||||
|
run: |
|
||||||
|
if [ "$(find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -l)" ]
|
||||||
|
then
|
||||||
|
find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -d
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
- name: Vet
|
- name: Vet
|
||||||
run: make vet
|
run: make vet
|
||||||
@@ -36,97 +44,107 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
version: v2.5
|
version: v2.5
|
||||||
|
|
||||||
- name: Test
|
|
||||||
run: make test
|
|
||||||
|
|
||||||
- name: End 2 end
|
|
||||||
run: make e2evv
|
|
||||||
|
|
||||||
- name: Build test mobile
|
|
||||||
run: make build-test-mobile
|
|
||||||
|
|
||||||
- uses: actions/upload-artifact@v6
|
|
||||||
with:
|
|
||||||
name: e2e packet flow linux-latest
|
|
||||||
path: e2e/mermaid/linux-latest
|
|
||||||
if-no-files-found: warn
|
|
||||||
|
|
||||||
test-linux-boringcrypto:
|
|
||||||
name: Build and test on linux with boringcrypto
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
steps:
|
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
|
||||||
with:
|
|
||||||
go-version: '1.25'
|
|
||||||
check-latest: true
|
|
||||||
|
|
||||||
- name: Build
|
|
||||||
run: make bin-boringcrypto
|
|
||||||
|
|
||||||
- name: Test
|
|
||||||
run: make test-boringcrypto
|
|
||||||
|
|
||||||
- name: End 2 end
|
|
||||||
run: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0"
|
|
||||||
|
|
||||||
test-linux-pkcs11:
|
|
||||||
name: Build and test on linux with pkcs11
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
steps:
|
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
|
||||||
with:
|
|
||||||
go-version: '1.25'
|
|
||||||
check-latest: true
|
|
||||||
|
|
||||||
- name: Build
|
|
||||||
run: make bin-pkcs11
|
|
||||||
|
|
||||||
- name: Test
|
|
||||||
run: make test-pkcs11
|
|
||||||
|
|
||||||
test:
|
test:
|
||||||
name: Build and test on ${{ matrix.os }}
|
name: Test ${{ matrix.name }}
|
||||||
runs-on: ${{ matrix.os }}
|
runs-on: ${{ matrix.os }}
|
||||||
strategy:
|
strategy:
|
||||||
|
fail-fast: false
|
||||||
matrix:
|
matrix:
|
||||||
os: [windows-latest, macos-latest]
|
include:
|
||||||
|
- name: linux
|
||||||
|
os: ubuntu-latest
|
||||||
|
build-cmd: go build ./cmd/nebula ./cmd/nebula-cert
|
||||||
|
test-cmd: make test
|
||||||
|
e2e-cmd: make e2evv
|
||||||
|
- name: linux-boringcrypto
|
||||||
|
os: ubuntu-latest
|
||||||
|
build-cmd: make bin-boringcrypto
|
||||||
|
test-cmd: make test-boringcrypto
|
||||||
|
e2e-cmd: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0"
|
||||||
|
- name: linux-pkcs11
|
||||||
|
os: ubuntu-latest
|
||||||
|
build-cmd: make bin-pkcs11
|
||||||
|
test-cmd: make test-pkcs11
|
||||||
|
e2e-cmd: ''
|
||||||
|
- name: macos
|
||||||
|
os: macos-latest
|
||||||
|
build-cmd: go build ./cmd/nebula ./cmd/nebula-cert
|
||||||
|
test-cmd: make test
|
||||||
|
e2e-cmd: make e2evv
|
||||||
|
- name: windows
|
||||||
|
os: windows-latest
|
||||||
|
build-cmd: go build ./cmd/nebula ./cmd/nebula-cert
|
||||||
|
test-cmd: make test
|
||||||
|
e2e-cmd: make e2evv
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.25'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build nebula
|
- name: Build
|
||||||
run: go build ./cmd/nebula
|
run: ${{ matrix.build-cmd }}
|
||||||
|
|
||||||
- name: Build nebula-cert
|
- name: Cross-build darwin-amd64
|
||||||
run: go build ./cmd/nebula-cert
|
if: matrix.name == 'macos'
|
||||||
|
run: GOARCH=amd64 go build -o /tmp/nebula-amd64 ./cmd/nebula && GOARCH=amd64 go build -o /tmp/nebula-cert-amd64 ./cmd/nebula-cert
|
||||||
- name: Vet
|
|
||||||
run: make vet
|
|
||||||
|
|
||||||
- name: golangci-lint
|
|
||||||
uses: golangci/golangci-lint-action@v9
|
|
||||||
with:
|
|
||||||
version: v2.5
|
|
||||||
|
|
||||||
- name: Test
|
- name: Test
|
||||||
run: make test
|
run: ${{ matrix.test-cmd }}
|
||||||
|
|
||||||
- name: End 2 end
|
- name: End 2 end
|
||||||
run: make e2evv
|
if: matrix.e2e-cmd != ''
|
||||||
|
run: ${{ matrix.e2e-cmd }}
|
||||||
|
|
||||||
- uses: actions/upload-artifact@v6
|
- uses: actions/upload-artifact@v7
|
||||||
|
if: matrix.e2e-cmd != '' && always()
|
||||||
with:
|
with:
|
||||||
name: e2e packet flow ${{ matrix.os }}
|
name: e2e packet flow ${{ matrix.name }}
|
||||||
path: e2e/mermaid/${{ matrix.os }}
|
path: e2e/mermaid/
|
||||||
if-no-files-found: warn
|
if-no-files-found: warn
|
||||||
|
|
||||||
|
cross-build:
|
||||||
|
name: Cross-build ${{ matrix.name }}
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
strategy:
|
||||||
|
fail-fast: false
|
||||||
|
matrix:
|
||||||
|
include:
|
||||||
|
- {name: linux-arm, make-target: all-cross-linux-arm}
|
||||||
|
- {name: linux-mips, make-target: all-cross-linux-mips}
|
||||||
|
- {name: linux-other, make-target: all-cross-linux-other}
|
||||||
|
- {name: freebsd, make-target: all-freebsd}
|
||||||
|
- {name: openbsd, make-target: all-openbsd}
|
||||||
|
- {name: netbsd, make-target: all-netbsd}
|
||||||
|
- {name: windows, make-target: all-cross-windows}
|
||||||
|
- {name: mobile, make-target: build-test-mobile}
|
||||||
|
steps:
|
||||||
|
|
||||||
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
|
- uses: actions/setup-go@v6
|
||||||
|
with:
|
||||||
|
go-version: '1.25'
|
||||||
|
check-latest: true
|
||||||
|
|
||||||
|
- name: Build ${{ matrix.name }}
|
||||||
|
run: make -j"$(nproc)" ${{ matrix.make-target }}
|
||||||
|
|
||||||
|
finish:
|
||||||
|
name: CI status
|
||||||
|
if: always()
|
||||||
|
needs: [static, test, cross-build]
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
|
||||||
|
- name: Fail if any upstream job failed
|
||||||
|
if: contains(needs.*.result, 'failure') || contains(needs.*.result, 'cancelled')
|
||||||
|
run: |
|
||||||
|
echo "upstream results: ${{ toJSON(needs) }}"
|
||||||
|
exit 1
|
||||||
|
|
||||||
|
- name: All upstream jobs passed
|
||||||
|
run: echo "ok"
|
||||||
|
|||||||
@@ -60,6 +60,18 @@ ALL = $(ALL_LINUX) \
|
|||||||
windows-amd64 \
|
windows-amd64 \
|
||||||
windows-arm64
|
windows-arm64
|
||||||
|
|
||||||
|
# Cross-build shards used by .github/workflows/test.yml — same as ALL_*
|
||||||
|
# but with the arch that has a native CI runner removed, so the cross-build
|
||||||
|
# job is not duplicating coverage the native test jobs already give.
|
||||||
|
ALL_CROSS_LINUX = $(filter-out linux-amd64,$(ALL_LINUX))
|
||||||
|
|
||||||
|
# ALL_CROSS_LINUX further split into family sub-shards so each can run on
|
||||||
|
# its own CI runner in parallel. Union of the three must equal
|
||||||
|
# ALL_CROSS_LINUX; adding a new linux arch goes into the matching family.
|
||||||
|
ALL_CROSS_LINUX_ARM = linux-arm-5 linux-arm-6 linux-arm-7 linux-arm64
|
||||||
|
ALL_CROSS_LINUX_MIPS = linux-mips linux-mipsle linux-mips64 linux-mips64le linux-mips-softfloat
|
||||||
|
ALL_CROSS_LINUX_OTHER = linux-386 linux-ppc64le linux-riscv64 linux-loong64
|
||||||
|
|
||||||
e2e:
|
e2e:
|
||||||
$(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e
|
$(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e
|
||||||
|
|
||||||
@@ -82,6 +94,35 @@ DOCKER_BIN = build/linux-amd64/nebula build/linux-amd64/nebula-cert
|
|||||||
|
|
||||||
all: $(ALL:%=build/%/nebula) $(ALL:%=build/%/nebula-cert)
|
all: $(ALL:%=build/%/nebula) $(ALL:%=build/%/nebula-cert)
|
||||||
|
|
||||||
|
all-linux: $(ALL_LINUX:%=build/%/nebula) $(ALL_LINUX:%=build/%/nebula-cert)
|
||||||
|
|
||||||
|
all-freebsd: $(ALL_FREEBSD:%=build/%/nebula) $(ALL_FREEBSD:%=build/%/nebula-cert)
|
||||||
|
|
||||||
|
all-openbsd: $(ALL_OPENBSD:%=build/%/nebula) $(ALL_OPENBSD:%=build/%/nebula-cert)
|
||||||
|
|
||||||
|
all-netbsd: $(ALL_NETBSD:%=build/%/nebula) $(ALL_NETBSD:%=build/%/nebula-cert)
|
||||||
|
|
||||||
|
all-darwin: build/darwin-amd64/nebula build/darwin-amd64/nebula-cert build/darwin-arm64/nebula build/darwin-arm64/nebula-cert
|
||||||
|
|
||||||
|
all-windows: build/windows-amd64/nebula.exe build/windows-amd64/nebula-cert.exe build/windows-arm64/nebula.exe build/windows-arm64/nebula-cert.exe
|
||||||
|
|
||||||
|
# CI cross-build shards. darwin-arm64 is covered by the native macos-latest
|
||||||
|
# job; windows-amd64 is covered by the native windows-latest job; both are
|
||||||
|
# omitted here to avoid building them a second time. darwin-amd64 stays in
|
||||||
|
# all-cross-darwin because intel mac is only a labeled/master-time native
|
||||||
|
# job, so PRs still need cross-build coverage for it.
|
||||||
|
all-cross-linux: $(ALL_CROSS_LINUX:%=build/%/nebula) $(ALL_CROSS_LINUX:%=build/%/nebula-cert)
|
||||||
|
|
||||||
|
all-cross-linux-arm: $(ALL_CROSS_LINUX_ARM:%=build/%/nebula) $(ALL_CROSS_LINUX_ARM:%=build/%/nebula-cert)
|
||||||
|
|
||||||
|
all-cross-linux-mips: $(ALL_CROSS_LINUX_MIPS:%=build/%/nebula) $(ALL_CROSS_LINUX_MIPS:%=build/%/nebula-cert)
|
||||||
|
|
||||||
|
all-cross-linux-other: $(ALL_CROSS_LINUX_OTHER:%=build/%/nebula) $(ALL_CROSS_LINUX_OTHER:%=build/%/nebula-cert)
|
||||||
|
|
||||||
|
all-cross-darwin: build/darwin-amd64/nebula build/darwin-amd64/nebula-cert
|
||||||
|
|
||||||
|
all-cross-windows: build/windows-arm64/nebula.exe build/windows-arm64/nebula-cert.exe
|
||||||
|
|
||||||
docker: docker/linux-$(shell go env GOARCH)
|
docker: docker/linux-$(shell go env GOARCH)
|
||||||
|
|
||||||
release: $(ALL:%=build/nebula-%.tar.gz)
|
release: $(ALL:%=build/nebula-%.tar.gz)
|
||||||
@@ -120,6 +161,10 @@ bin-pkcs11: BUILD_ARGS += -tags pkcs11
|
|||||||
bin-pkcs11: CGO_ENABLED = 1
|
bin-pkcs11: CGO_ENABLED = 1
|
||||||
bin-pkcs11: bin
|
bin-pkcs11: bin
|
||||||
|
|
||||||
|
# Build with the pprof debug server (serves on :6060). See startPprofServer.
|
||||||
|
debug: BUILD_ARGS += -tags debug
|
||||||
|
debug: bin
|
||||||
|
|
||||||
bin:
|
bin:
|
||||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
||||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
||||||
@@ -227,6 +272,9 @@ smoke-relay-docker: bin-docker
|
|||||||
cd .github/workflows/smoke/ && ./build-relay.sh
|
cd .github/workflows/smoke/ && ./build-relay.sh
|
||||||
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
||||||
|
|
||||||
|
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
|
||||||
|
smoke-docker-ipv6: smoke-docker
|
||||||
|
|
||||||
smoke-docker-race: BUILD_ARGS = -race
|
smoke-docker-race: BUILD_ARGS = -race
|
||||||
smoke-docker-race: CGO_ENABLED = 1
|
smoke-docker-race: CGO_ENABLED = 1
|
||||||
smoke-docker-race: smoke-docker
|
smoke-docker-race: smoke-docker
|
||||||
@@ -236,5 +284,5 @@ smoke-vagrant/%: bin-docker build/%/nebula
|
|||||||
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
||||||
|
|
||||||
.FORCE:
|
.FORCE:
|
||||||
.PHONY: bench bench-cpu bench-cpu-long bin build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin debug build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
||||||
.DEFAULT_GOAL := bin
|
.DEFAULT_GOAL := bin
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
+10
-4
@@ -13,6 +13,12 @@ import (
|
|||||||
"golang.org/x/crypto/ed25519"
|
"golang.org/x/crypto/ed25519"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// testCertNow is the reference "now" used to derive default before/after times
|
||||||
|
// in NewTestCaCert and NewTestCert. Holding it fixed for the lifetime of the
|
||||||
|
// test binary keeps CA and leaf defaults aligned at the same second, so a leaf
|
||||||
|
// signed with default times can never expire after its CA on a rounding race.
|
||||||
|
var testCertNow = time.Now().Round(time.Second)
|
||||||
|
|
||||||
// NewTestCaCert will create a new ca certificate
|
// NewTestCaCert will create a new ca certificate
|
||||||
func NewTestCaCert(version Version, curve Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
|
func NewTestCaCert(version Version, curve Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
|
||||||
var err error
|
var err error
|
||||||
@@ -34,10 +40,10 @@ func NewTestCaCert(version Version, curve Curve, before, after time.Time, networ
|
|||||||
}
|
}
|
||||||
|
|
||||||
if before.IsZero() {
|
if before.IsZero() {
|
||||||
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
before = testCertNow.Add(time.Second * -60)
|
||||||
}
|
}
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
after = testCertNow.Add(time.Second * 60)
|
||||||
}
|
}
|
||||||
|
|
||||||
t := &TBSCertificate{
|
t := &TBSCertificate{
|
||||||
@@ -70,11 +76,11 @@ func NewTestCaCert(version Version, curve Curve, before, after time.Time, networ
|
|||||||
// Expiry times are defaulted if you do not pass them in
|
// Expiry times are defaulted if you do not pass them in
|
||||||
func NewTestCert(v Version, curve Curve, ca Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
|
func NewTestCert(v Version, curve Curve, ca Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
|
||||||
if before.IsZero() {
|
if before.IsZero() {
|
||||||
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
before = testCertNow.Add(time.Second * -60)
|
||||||
}
|
}
|
||||||
|
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
after = testCertNow.Add(time.Second * 60)
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(networks) == 0 {
|
if len(networks) == 0 {
|
||||||
|
|||||||
+32
-2
@@ -148,6 +148,9 @@ func MarshalSigningPublicKeyToPEM(curve Curve, b []byte) []byte {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// UnmarshalPublicKeyFromPEM will try to unmarshal the first pem block in a byte array, returning any non
|
||||||
|
// consumed data or an error on failure. Only key-agreement (ECDH) public key banners are accepted.
|
||||||
|
// Use UnmarshalSigningPublicKeyFromPEM for Ed25519/ECDSA banners.
|
||||||
func UnmarshalPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) {
|
func UnmarshalPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) {
|
||||||
k, r := pem.Decode(b)
|
k, r := pem.Decode(b)
|
||||||
if k == nil {
|
if k == nil {
|
||||||
@@ -156,10 +159,10 @@ func UnmarshalPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) {
|
|||||||
var expectedLen int
|
var expectedLen int
|
||||||
var curve Curve
|
var curve Curve
|
||||||
switch k.Type {
|
switch k.Type {
|
||||||
case X25519PublicKeyBanner, Ed25519PublicKeyBanner:
|
case X25519PublicKeyBanner:
|
||||||
expectedLen = 32
|
expectedLen = 32
|
||||||
curve = Curve_CURVE25519
|
curve = Curve_CURVE25519
|
||||||
case P256PublicKeyBanner, ECDSAP256PublicKeyBanner:
|
case P256PublicKeyBanner:
|
||||||
// Uncompressed
|
// Uncompressed
|
||||||
expectedLen = 65
|
expectedLen = 65
|
||||||
curve = Curve_P256
|
curve = Curve_P256
|
||||||
@@ -172,6 +175,33 @@ func UnmarshalPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) {
|
|||||||
return k.Bytes, r, curve, nil
|
return k.Bytes, r, curve, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// UnmarshalSigningPublicKeyFromPEM will try to unmarshal the first pem block in a byte array, returning any non
|
||||||
|
// consumed data or an error on failure. Only Ed25519/ECDSA public key banners are accepted.
|
||||||
|
// Use UnmarshalPublicKeyFromPEM for X25519/P256 (ECDH) banners.
|
||||||
|
func UnmarshalSigningPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) {
|
||||||
|
k, r := pem.Decode(b)
|
||||||
|
if k == nil {
|
||||||
|
return nil, r, 0, fmt.Errorf("input did not contain a valid PEM encoded block")
|
||||||
|
}
|
||||||
|
var expectedLen int
|
||||||
|
var curve Curve
|
||||||
|
switch k.Type {
|
||||||
|
case Ed25519PublicKeyBanner:
|
||||||
|
expectedLen = 32
|
||||||
|
curve = Curve_CURVE25519
|
||||||
|
case ECDSAP256PublicKeyBanner:
|
||||||
|
// Uncompressed
|
||||||
|
expectedLen = 65
|
||||||
|
curve = Curve_P256
|
||||||
|
default:
|
||||||
|
return nil, r, 0, fmt.Errorf("bytes did not contain a proper Ed25519/ECDSA public key banner")
|
||||||
|
}
|
||||||
|
if len(k.Bytes) != expectedLen {
|
||||||
|
return nil, r, 0, fmt.Errorf("key was not %d bytes, is invalid %s public key", expectedLen, curve)
|
||||||
|
}
|
||||||
|
return k.Bytes, r, curve, nil
|
||||||
|
}
|
||||||
|
|
||||||
func MarshalPrivateKeyToPEM(curve Curve, b []byte) []byte {
|
func MarshalPrivateKeyToPEM(curve Curve, b []byte) []byte {
|
||||||
switch curve {
|
switch curve {
|
||||||
case Curve_CURVE25519:
|
case Curve_CURVE25519:
|
||||||
|
|||||||
+88
-68
@@ -255,60 +255,6 @@ AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|||||||
func TestUnmarshalPublicKeyFromPEM(t *testing.T) {
|
func TestUnmarshalPublicKeyFromPEM(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
pubKey := []byte(`# A good key
|
pubKey := []byte(`# A good key
|
||||||
-----BEGIN NEBULA ED25519 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA ED25519 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
shortKey := []byte(`# A short key
|
|
||||||
-----BEGIN NEBULA ED25519 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==
|
|
||||||
-----END NEBULA ED25519 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
invalidBanner := []byte(`# Invalid banner
|
|
||||||
-----BEGIN NOT A NEBULA PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NOT A NEBULA PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
invalidPem := []byte(`# Not a valid PEM format
|
|
||||||
-BEGIN NEBULA ED25519 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-END NEBULA ED25519 PUBLIC KEY-----`)
|
|
||||||
|
|
||||||
keyBundle := appendByteSlices(pubKey, shortKey, invalidBanner, invalidPem)
|
|
||||||
|
|
||||||
// Success test case
|
|
||||||
k, rest, curve, err := UnmarshalPublicKeyFromPEM(keyBundle)
|
|
||||||
assert.Len(t, k, 32)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
|
|
||||||
|
|
||||||
// Fail due to short key
|
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem))
|
|
||||||
require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 public key")
|
|
||||||
|
|
||||||
// Fail due to invalid banner
|
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
require.EqualError(t, err, "bytes did not contain a proper public key banner")
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
|
|
||||||
// Fail due to invalid PEM format, because
|
|
||||||
// it's missing the requisite pre-encapsulation boundary.
|
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
require.EqualError(t, err, "input did not contain a valid PEM encoded block")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUnmarshalX25519PublicKey(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
pubKey := []byte(`# A good key
|
|
||||||
-----BEGIN NEBULA X25519 PUBLIC KEY-----
|
-----BEGIN NEBULA X25519 PUBLIC KEY-----
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
-----END NEBULA X25519 PUBLIC KEY-----
|
-----END NEBULA X25519 PUBLIC KEY-----
|
||||||
@@ -319,7 +265,7 @@ AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
|||||||
AAAAAAAAAAAAAAAAAAAAAAA=
|
AAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
-----END NEBULA P256 PUBLIC KEY-----
|
-----END NEBULA P256 PUBLIC KEY-----
|
||||||
`)
|
`)
|
||||||
oldPubP256Key := []byte(`# A good key
|
signingKey := []byte(`# A signing key has the wrong scope for this function
|
||||||
-----BEGIN NEBULA ECDSA P256 PUBLIC KEY-----
|
-----BEGIN NEBULA ECDSA P256 PUBLIC KEY-----
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
||||||
AAAAAAAAAAAAAAAAAAAAAAA=
|
AAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
@@ -340,44 +286,118 @@ AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
-END NEBULA X25519 PUBLIC KEY-----`)
|
-END NEBULA X25519 PUBLIC KEY-----`)
|
||||||
|
|
||||||
keyBundle := appendByteSlices(pubKey, pubP256Key, oldPubP256Key, shortKey, invalidBanner, invalidPem)
|
keyBundle := appendByteSlices(pubKey, pubP256Key, signingKey, shortKey, invalidBanner, invalidPem)
|
||||||
|
|
||||||
// Success test case
|
// X25519 key
|
||||||
k, rest, curve, err := UnmarshalPublicKeyFromPEM(keyBundle)
|
k, rest, curve, err := UnmarshalPublicKeyFromPEM(keyBundle)
|
||||||
assert.Len(t, k, 32)
|
assert.Len(t, k, 32)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, rest, appendByteSlices(pubP256Key, oldPubP256Key, shortKey, invalidBanner, invalidPem))
|
assert.Equal(t, rest, appendByteSlices(pubP256Key, signingKey, shortKey, invalidBanner, invalidPem))
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
assert.Equal(t, Curve_CURVE25519, curve)
|
||||||
|
|
||||||
// Success test case
|
// P256 key
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
||||||
assert.Len(t, k, 65)
|
assert.Len(t, k, 65)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, rest, appendByteSlices(oldPubP256Key, shortKey, invalidBanner, invalidPem))
|
assert.Equal(t, rest, appendByteSlices(signingKey, shortKey, invalidBanner, invalidPem))
|
||||||
assert.Equal(t, Curve_P256, curve)
|
assert.Equal(t, Curve_P256, curve)
|
||||||
|
|
||||||
// Success test case
|
// Reject a signing public key (Ed25519/ECDSA banner)
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
k, rest, _, err = UnmarshalPublicKeyFromPEM(rest)
|
||||||
assert.Len(t, k, 65)
|
assert.Nil(t, k)
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
|
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
|
||||||
assert.Equal(t, Curve_P256, curve)
|
require.EqualError(t, err, "bytes did not contain a proper public key banner")
|
||||||
|
|
||||||
// Fail due to short key
|
// Fail due to short key
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
k, rest, _, err = UnmarshalPublicKeyFromPEM(rest)
|
||||||
assert.Nil(t, k)
|
assert.Nil(t, k)
|
||||||
assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem))
|
assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem))
|
||||||
require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 public key")
|
require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 public key")
|
||||||
|
|
||||||
// Fail due to invalid banner
|
// Fail due to invalid banner
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
k, rest, _, err = UnmarshalPublicKeyFromPEM(rest)
|
||||||
assert.Nil(t, k)
|
assert.Nil(t, k)
|
||||||
require.EqualError(t, err, "bytes did not contain a proper public key banner")
|
require.EqualError(t, err, "bytes did not contain a proper public key banner")
|
||||||
assert.Equal(t, rest, invalidPem)
|
assert.Equal(t, rest, invalidPem)
|
||||||
|
|
||||||
// Fail due to invalid PEM format, because
|
// Fail due to invalid PEM format, because
|
||||||
// it's missing the requisite pre-encapsulation boundary.
|
// it's missing the requisite pre-encapsulation boundary.
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
k, rest, _, err = UnmarshalPublicKeyFromPEM(rest)
|
||||||
|
assert.Nil(t, k)
|
||||||
|
assert.Equal(t, rest, invalidPem)
|
||||||
|
require.EqualError(t, err, "input did not contain a valid PEM encoded block")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUnmarshalSigningPublicKeyFromPEM(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
pubKey := []byte(`# A good key
|
||||||
|
-----BEGIN NEBULA ED25519 PUBLIC KEY-----
|
||||||
|
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
|
-----END NEBULA ED25519 PUBLIC KEY-----
|
||||||
|
`)
|
||||||
|
pubP256Key := []byte(`# A good key
|
||||||
|
-----BEGIN NEBULA ECDSA P256 PUBLIC KEY-----
|
||||||
|
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
||||||
|
AAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
|
-----END NEBULA ECDSA P256 PUBLIC KEY-----
|
||||||
|
`)
|
||||||
|
ecdhKey := []byte(`# A key-agreement key has the wrong scope for this function
|
||||||
|
-----BEGIN NEBULA X25519 PUBLIC KEY-----
|
||||||
|
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
|
-----END NEBULA X25519 PUBLIC KEY-----
|
||||||
|
`)
|
||||||
|
shortKey := []byte(`# A short key
|
||||||
|
-----BEGIN NEBULA ED25519 PUBLIC KEY-----
|
||||||
|
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==
|
||||||
|
-----END NEBULA ED25519 PUBLIC KEY-----
|
||||||
|
`)
|
||||||
|
invalidBanner := []byte(`# Invalid banner
|
||||||
|
-----BEGIN NOT A NEBULA PUBLIC KEY-----
|
||||||
|
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
|
-----END NOT A NEBULA PUBLIC KEY-----
|
||||||
|
`)
|
||||||
|
invalidPem := []byte(`# Not a valid PEM format
|
||||||
|
-BEGIN NEBULA ED25519 PUBLIC KEY-----
|
||||||
|
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
|
-END NEBULA ED25519 PUBLIC KEY-----`)
|
||||||
|
|
||||||
|
keyBundle := appendByteSlices(pubKey, pubP256Key, ecdhKey, shortKey, invalidBanner, invalidPem)
|
||||||
|
|
||||||
|
// Ed25519 key
|
||||||
|
k, rest, curve, err := UnmarshalSigningPublicKeyFromPEM(keyBundle)
|
||||||
|
assert.Len(t, k, 32)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, rest, appendByteSlices(pubP256Key, ecdhKey, shortKey, invalidBanner, invalidPem))
|
||||||
|
assert.Equal(t, Curve_CURVE25519, curve)
|
||||||
|
|
||||||
|
// ECDSA P256 key
|
||||||
|
k, rest, curve, err = UnmarshalSigningPublicKeyFromPEM(rest)
|
||||||
|
assert.Len(t, k, 65)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, rest, appendByteSlices(ecdhKey, shortKey, invalidBanner, invalidPem))
|
||||||
|
assert.Equal(t, Curve_P256, curve)
|
||||||
|
|
||||||
|
// Reject a key-agreement public key (X25519/P256 banner)
|
||||||
|
k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest)
|
||||||
|
assert.Nil(t, k)
|
||||||
|
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
|
||||||
|
require.EqualError(t, err, "bytes did not contain a proper Ed25519/ECDSA public key banner")
|
||||||
|
|
||||||
|
// Fail due to short key
|
||||||
|
k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest)
|
||||||
|
assert.Nil(t, k)
|
||||||
|
assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem))
|
||||||
|
require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 public key")
|
||||||
|
|
||||||
|
// Fail due to invalid banner
|
||||||
|
k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest)
|
||||||
|
assert.Nil(t, k)
|
||||||
|
require.EqualError(t, err, "bytes did not contain a proper Ed25519/ECDSA public key banner")
|
||||||
|
assert.Equal(t, rest, invalidPem)
|
||||||
|
|
||||||
|
// Fail due to invalid PEM format, because
|
||||||
|
// it's missing the requisite pre-encapsulation boundary.
|
||||||
|
k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest)
|
||||||
assert.Nil(t, k)
|
assert.Nil(t, k)
|
||||||
assert.Equal(t, rest, invalidPem)
|
assert.Equal(t, rest, invalidPem)
|
||||||
require.EqualError(t, err, "input did not contain a valid PEM encoded block")
|
require.EqualError(t, err, "input did not contain a valid PEM encoded block")
|
||||||
|
|||||||
+10
-4
@@ -14,6 +14,12 @@ import (
|
|||||||
"golang.org/x/crypto/ed25519"
|
"golang.org/x/crypto/ed25519"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// testCertNow is the reference "now" used to derive default before/after times
|
||||||
|
// in NewTestCaCert and NewTestCert. Holding it fixed for the lifetime of the
|
||||||
|
// test binary keeps CA and leaf defaults aligned at the same second, so a leaf
|
||||||
|
// signed with default times can never expire after its CA on a rounding race.
|
||||||
|
var testCertNow = time.Now().Round(time.Second)
|
||||||
|
|
||||||
// NewTestCaCert will create a new ca certificate
|
// NewTestCaCert will create a new ca certificate
|
||||||
func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
|
func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
|
||||||
var err error
|
var err error
|
||||||
@@ -35,10 +41,10 @@ func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Ti
|
|||||||
}
|
}
|
||||||
|
|
||||||
if before.IsZero() {
|
if before.IsZero() {
|
||||||
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
before = testCertNow.Add(time.Second * -60)
|
||||||
}
|
}
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
after = testCertNow.Add(time.Second * 60)
|
||||||
}
|
}
|
||||||
|
|
||||||
t := &cert.TBSCertificate{
|
t := &cert.TBSCertificate{
|
||||||
@@ -71,11 +77,11 @@ func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Ti
|
|||||||
// Expiry times are defaulted if you do not pass them in
|
// Expiry times are defaulted if you do not pass them in
|
||||||
func NewTestCert(v cert.Version, curve cert.Curve, ca cert.Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
|
func NewTestCert(v cert.Version, curve cert.Curve, ca cert.Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
|
||||||
if before.IsZero() {
|
if before.IsZero() {
|
||||||
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
before = testCertNow.Add(time.Second * -60)
|
||||||
}
|
}
|
||||||
|
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
after = testCertNow.Add(time.Second * 60)
|
||||||
}
|
}
|
||||||
|
|
||||||
var pub, priv []byte
|
var pub, priv []byte
|
||||||
|
|||||||
+32
-7
@@ -97,6 +97,19 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
|
if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
} else {
|
||||||
|
// out-key is meaningless under PKCS#11 because the private key never
|
||||||
|
// leaves the HSM; reject it so we never silently accept or claim a
|
||||||
|
// stdout slot for it.
|
||||||
|
outKeySet := false
|
||||||
|
cf.set.Visit(func(f *flag.Flag) {
|
||||||
|
if f.Name == "out-key" {
|
||||||
|
outKeySet = true
|
||||||
|
}
|
||||||
|
})
|
||||||
|
if outKeySet {
|
||||||
|
return newHelpErrorf("cannot set -out-key with -pkcs11")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if err := mustFlagString("out-crt", cf.outCertPath); err != nil {
|
if err := mustFlagString("out-crt", cf.outCertPath); err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -171,12 +184,21 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var claims ioClaims
|
||||||
|
if err := reserveOutputs(&claims,
|
||||||
|
"out-key", *cf.outKeyPath,
|
||||||
|
"out-crt", *cf.outCertPath,
|
||||||
|
"out-qr", *cf.outQRPath,
|
||||||
|
); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
var passphrase []byte
|
var passphrase []byte
|
||||||
if !isP11 && *cf.encryption {
|
if !isP11 && *cf.encryption {
|
||||||
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
|
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
|
||||||
if len(passphrase) == 0 {
|
if len(passphrase) == 0 {
|
||||||
for i := 0; i < 5; i++ {
|
for i := 0; i < 5; i++ {
|
||||||
out.Write([]byte("Enter passphrase: "))
|
errOut.Write([]byte("Enter passphrase: "))
|
||||||
passphrase, err = pr.ReadPassword()
|
passphrase, err = pr.ReadPassword()
|
||||||
|
|
||||||
if err == ErrNoTerminal {
|
if err == ErrNoTerminal {
|
||||||
@@ -261,14 +283,16 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
Curve: curve,
|
Curve: curve,
|
||||||
}
|
}
|
||||||
|
|
||||||
if !isP11 {
|
if !isP11 && !isStdio(*cf.outKeyPath) {
|
||||||
if _, err := os.Stat(*cf.outKeyPath); err == nil {
|
if _, err := os.Stat(*cf.outKeyPath); err == nil {
|
||||||
return fmt.Errorf("refusing to overwrite existing CA key: %s", *cf.outKeyPath)
|
return fmt.Errorf("refusing to overwrite existing CA key: %s", *cf.outKeyPath)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := os.Stat(*cf.outCertPath); err == nil {
|
if !isStdio(*cf.outCertPath) {
|
||||||
return fmt.Errorf("refusing to overwrite existing CA cert: %s", *cf.outCertPath)
|
if _, err := os.Stat(*cf.outCertPath); err == nil {
|
||||||
|
return fmt.Errorf("refusing to overwrite existing CA cert: %s", *cf.outCertPath)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var c cert.Certificate
|
var c cert.Certificate
|
||||||
@@ -294,7 +318,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
b = cert.MarshalSigningPrivateKeyToPEM(curve, rawPriv)
|
b = cert.MarshalSigningPrivateKeyToPEM(curve, rawPriv)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*cf.outKeyPath, b, 0600)
|
err = writeOutput(*cf.outKeyPath, b, 0600, out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-key: %s", err)
|
return fmt.Errorf("error while writing out-key: %s", err)
|
||||||
}
|
}
|
||||||
@@ -305,7 +329,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
return fmt.Errorf("error while marshalling certificate: %s", err)
|
return fmt.Errorf("error while marshalling certificate: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*cf.outCertPath, b, 0600)
|
err = writeOutput(*cf.outCertPath, b, 0600, out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-crt: %s", err)
|
return fmt.Errorf("error while writing out-crt: %s", err)
|
||||||
}
|
}
|
||||||
@@ -316,7 +340,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
return fmt.Errorf("error while generating qr code: %s", err)
|
return fmt.Errorf("error while generating qr code: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*cf.outQRPath, b, 0600)
|
err = writeOutput(*cf.outQRPath, b, 0600, out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-qr: %s", err)
|
return fmt.Errorf("error while writing out-qr: %s", err)
|
||||||
}
|
}
|
||||||
@@ -332,6 +356,7 @@ func caSummary() string {
|
|||||||
func caHelp(out io.Writer) {
|
func caHelp(out io.Writer) {
|
||||||
cf := newCaFlags()
|
cf := newCaFlags()
|
||||||
out.Write([]byte("Usage of " + os.Args[0] + " " + caSummary() + "\n"))
|
out.Write([]byte("Usage of " + os.Args[0] + " " + caSummary() + "\n"))
|
||||||
|
out.Write([]byte(stdioHelpText))
|
||||||
cf.set.SetOutput(out)
|
cf.set.SetOutput(out)
|
||||||
cf.set.PrintDefaults()
|
cf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ func Test_caHelp(t *testing.T) {
|
|||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
"Usage of "+os.Args[0]+" ca <flags>: create a self signed certificate authority\n"+
|
"Usage of "+os.Args[0]+" ca <flags>: create a self signed certificate authority\n"+
|
||||||
|
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
|
||||||
" -argon-iterations uint\n"+
|
" -argon-iterations uint\n"+
|
||||||
" \tOptional: Argon2 iterations parameter used for encrypted private key passphrase (default 1)\n"+
|
" \tOptional: Argon2 iterations parameter used for encrypted private key passphrase (default 1)\n"+
|
||||||
" -argon-memory uint\n"+
|
" -argon-memory uint\n"+
|
||||||
@@ -84,7 +85,7 @@ func Test_ca(t *testing.T) {
|
|||||||
err: nil,
|
err: nil,
|
||||||
}
|
}
|
||||||
|
|
||||||
pwPromptOb := "Enter passphrase: "
|
pwPromptEB := "Enter passphrase: "
|
||||||
|
|
||||||
// required args
|
// required args
|
||||||
assertHelpError(t, ca(
|
assertHelpError(t, ca(
|
||||||
@@ -168,8 +169,8 @@ func Test_ca(t *testing.T) {
|
|||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.NoError(t, ca(args, ob, eb, testpw))
|
require.NoError(t, ca(args, ob, eb, testpw))
|
||||||
assert.Equal(t, pwPromptOb, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, pwPromptEB, eb.String())
|
||||||
|
|
||||||
// test encrypted key with passphrase environment variable
|
// test encrypted key with passphrase environment variable
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
@@ -207,8 +208,8 @@ func Test_ca(t *testing.T) {
|
|||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.Error(t, ca(args, ob, eb, errpw))
|
require.Error(t, ca(args, ob, eb, errpw))
|
||||||
assert.Equal(t, pwPromptOb, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, pwPromptEB, eb.String())
|
||||||
|
|
||||||
// test when user fails to enter a password
|
// test when user fails to enter a password
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
@@ -217,8 +218,8 @@ func Test_ca(t *testing.T) {
|
|||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.EqualError(t, ca(args, ob, eb, nopw), "no passphrase specified, remove -encrypt flag to write out-key in plaintext")
|
require.EqualError(t, ca(args, ob, eb, nopw), "no passphrase specified, remove -encrypt flag to write out-key in plaintext")
|
||||||
assert.Equal(t, strings.Repeat(pwPromptOb, 5), ob.String()) // prompts 5 times before giving up
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, strings.Repeat(pwPromptEB, 5), eb.String()) // prompts 5 times before giving up
|
||||||
|
|
||||||
// create valid cert/key for overwrite tests
|
// create valid cert/key for overwrite tests
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
@@ -247,3 +248,67 @@ func Test_ca(t *testing.T) {
|
|||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func Test_ca_stdio(t *testing.T) {
|
||||||
|
nopw := &StubPasswordReader{}
|
||||||
|
|
||||||
|
keyF, err := os.CreateTemp("", "ca.key")
|
||||||
|
require.NoError(t, err)
|
||||||
|
os.Remove(keyF.Name())
|
||||||
|
defer os.Remove(keyF.Name())
|
||||||
|
|
||||||
|
crtF, err := os.CreateTemp("", "ca.crt")
|
||||||
|
require.NoError(t, err)
|
||||||
|
os.Remove(crtF.Name())
|
||||||
|
defer os.Remove(crtF.Name())
|
||||||
|
|
||||||
|
// out-crt on stdout, out-key on disk
|
||||||
|
ob := &bytes.Buffer{}
|
||||||
|
eb := &bytes.Buffer{}
|
||||||
|
require.NoError(t, ca([]string{"-name", "test-ca", "-duration", "1h", "-out-crt", "-", "-out-key", keyF.Name()}, ob, eb, nopw))
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
c, _, err := cert.UnmarshalCertificateFromPEM(ob.Bytes())
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.True(t, c.IsCA())
|
||||||
|
assert.Equal(t, "test-ca", c.Name())
|
||||||
|
|
||||||
|
// out-key on stdout, out-crt on disk
|
||||||
|
os.Remove(keyF.Name())
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
require.NoError(t, ca([]string{"-name", "test-ca", "-duration", "1h", "-out-crt", crtF.Name(), "-out-key", "-"}, ob, eb, nopw))
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
_, _, curve, err := cert.UnmarshalSigningPrivateKeyFromPEM(ob.Bytes())
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
||||||
|
|
||||||
|
// dual stdout is rejected up front
|
||||||
|
os.Remove(crtF.Name())
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
require.EqualError(t,
|
||||||
|
ca([]string{"-name", "test-ca", "-duration", "1h", "-out-crt", "-", "-out-key", "-"}, ob, eb, nopw),
|
||||||
|
`-out-key and -out-crt both set to "-", only one output may write to stdout`)
|
||||||
|
assert.Empty(t, ob.String())
|
||||||
|
|
||||||
|
// an output conflict combined with -encrypt must error BEFORE prompting
|
||||||
|
// for a passphrase; pr would record any read attempt
|
||||||
|
tracker := &trackingPasswordReader{}
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
require.EqualError(t,
|
||||||
|
ca([]string{"-name", "test-ca", "-duration", "1h", "-encrypt", "-out-crt", "-", "-out-key", "-"}, ob, eb, tracker),
|
||||||
|
`-out-key and -out-crt both set to "-", only one output may write to stdout`)
|
||||||
|
assert.Empty(t, ob.String())
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
assert.Zero(t, tracker.calls, "passphrase prompt should not have been called")
|
||||||
|
}
|
||||||
|
|
||||||
|
type trackingPasswordReader struct {
|
||||||
|
calls int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (pr *trackingPasswordReader) ReadPassword() ([]byte, error) {
|
||||||
|
pr.calls++
|
||||||
|
return []byte(""), nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -42,6 +42,8 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
|
if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
} else if *cf.outKeyPath != "" {
|
||||||
|
return newHelpErrorf("cannot set -out-key with -pkcs11")
|
||||||
}
|
}
|
||||||
if err = mustFlagString("out-pub", cf.outPubPath); err != nil {
|
if err = mustFlagString("out-pub", cf.outPubPath); err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -69,6 +71,14 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var claims ioClaims
|
||||||
|
if err := reserveOutputs(&claims,
|
||||||
|
"out-key", *cf.outKeyPath,
|
||||||
|
"out-pub", *cf.outPubPath,
|
||||||
|
); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
if isP11 {
|
if isP11 {
|
||||||
p11Client, err := pkclient.FromUrl(*cf.p11url)
|
p11Client, err := pkclient.FromUrl(*cf.p11url)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -82,12 +92,12 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return fmt.Errorf("error while getting public key: %w", err)
|
return fmt.Errorf("error while getting public key: %w", err)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
err = os.WriteFile(*cf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600)
|
err = writeOutput(*cf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600, out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-key: %s", err)
|
return fmt.Errorf("error while writing out-key: %s", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
err = os.WriteFile(*cf.outPubPath, cert.MarshalPublicKeyToPEM(curve, pub), 0600)
|
err = writeOutput(*cf.outPubPath, cert.MarshalPublicKeyToPEM(curve, pub), 0600, out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-pub: %s", err)
|
return fmt.Errorf("error while writing out-pub: %s", err)
|
||||||
}
|
}
|
||||||
@@ -102,6 +112,7 @@ func keygenSummary() string {
|
|||||||
func keygenHelp(out io.Writer) {
|
func keygenHelp(out io.Writer) {
|
||||||
cf := newKeygenFlags()
|
cf := newKeygenFlags()
|
||||||
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + keygenSummary() + "\n"))
|
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + keygenSummary() + "\n"))
|
||||||
|
_, _ = out.Write([]byte(stdioHelpText))
|
||||||
cf.set.SetOutput(out)
|
cf.set.SetOutput(out)
|
||||||
cf.set.PrintDefaults()
|
cf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ func Test_keygenHelp(t *testing.T) {
|
|||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
"Usage of "+os.Args[0]+" keygen <flags>: create a public/private key pair. the public key can be passed to `nebula-cert sign`\n"+
|
"Usage of "+os.Args[0]+" keygen <flags>: create a public/private key pair. the public key can be passed to `nebula-cert sign`\n"+
|
||||||
|
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
|
||||||
" -curve string\n"+
|
" -curve string\n"+
|
||||||
" \tECDH Curve (25519, P256) (default \"25519\")\n"+
|
" \tECDH Curve (25519, P256) (default \"25519\")\n"+
|
||||||
" -out-key string\n"+
|
" -out-key string\n"+
|
||||||
@@ -93,3 +94,43 @@ func Test_keygen(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Len(t, lPub, 32)
|
assert.Len(t, lPub, 32)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func Test_keygen_stdio(t *testing.T) {
|
||||||
|
keyF, err := os.CreateTemp("", "test.key")
|
||||||
|
require.NoError(t, err)
|
||||||
|
os.Remove(keyF.Name())
|
||||||
|
defer os.Remove(keyF.Name())
|
||||||
|
|
||||||
|
pubF, err := os.CreateTemp("", "test.pub")
|
||||||
|
require.NoError(t, err)
|
||||||
|
os.Remove(pubF.Name())
|
||||||
|
defer os.Remove(pubF.Name())
|
||||||
|
|
||||||
|
// out-pub on stdout, out-key on disk
|
||||||
|
ob := &bytes.Buffer{}
|
||||||
|
eb := &bytes.Buffer{}
|
||||||
|
require.NoError(t, keygen([]string{"-out-pub", "-", "-out-key", keyF.Name()}, ob, eb))
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
lPub, _, curve, err := cert.UnmarshalPublicKeyFromPEM(ob.Bytes())
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
||||||
|
assert.Len(t, lPub, 32)
|
||||||
|
|
||||||
|
// out-key on stdout, out-pub on disk
|
||||||
|
os.Remove(keyF.Name())
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
require.NoError(t, keygen([]string{"-out-pub", pubF.Name(), "-out-key", "-"}, ob, eb))
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
lKey, _, curve, err := cert.UnmarshalPrivateKeyFromPEM(ob.Bytes())
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
||||||
|
assert.Len(t, lKey, 32)
|
||||||
|
|
||||||
|
// both on stdout is a conflict caught up front
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
require.EqualError(t, keygen([]string{"-out-pub", "-", "-out-key", "-"}, ob, eb),
|
||||||
|
`-out-key and -out-pub both set to "-", only one output may write to stdout`)
|
||||||
|
assert.Empty(t, ob.String())
|
||||||
|
}
|
||||||
|
|||||||
@@ -22,7 +22,9 @@ func (pr StdinPasswordReader) ReadPassword() ([]byte, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
password, err := term.ReadPassword(int(os.Stdin.Fd()))
|
password, err := term.ReadPassword(int(os.Stdin.Fd()))
|
||||||
fmt.Println()
|
// Terminal echo is off while reading, so the user's Enter key does not
|
||||||
|
// produce a visible newline. Emit one on stderr to match the prompt.
|
||||||
|
fmt.Fprintln(os.Stderr)
|
||||||
|
|
||||||
return password, err
|
return password, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -40,11 +40,23 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
rawCert, err := os.ReadFile(*pf.path)
|
var claims ioClaims
|
||||||
|
if err := reserveInputs(&claims, "path", *pf.path); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := reserveOutputs(&claims, "out-qr", *pf.outQRPath); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
rawCert, err := readInput("path", *pf.path, &claims)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("unable to read cert; %s", err)
|
return fmt.Errorf("unable to read cert; %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// When the QR is going to stdout, suppress the human-readable text/json
|
||||||
|
// output so the binary stream is not contaminated.
|
||||||
|
qrToStdout := isStdio(*pf.outQRPath)
|
||||||
|
|
||||||
var c cert.Certificate
|
var c cert.Certificate
|
||||||
var qrBytes []byte
|
var qrBytes []byte
|
||||||
part := 0
|
part := 0
|
||||||
@@ -57,11 +69,13 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return fmt.Errorf("error while unmarshaling cert: %s", err)
|
return fmt.Errorf("error while unmarshaling cert: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if *pf.json {
|
if !qrToStdout {
|
||||||
jsonCerts = append(jsonCerts, c)
|
if *pf.json {
|
||||||
} else {
|
jsonCerts = append(jsonCerts, c)
|
||||||
_, _ = out.Write([]byte(c.String()))
|
} else {
|
||||||
_, _ = out.Write([]byte("\n"))
|
_, _ = out.Write([]byte(c.String()))
|
||||||
|
_, _ = out.Write([]byte("\n"))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if *pf.outQRPath != "" {
|
if *pf.outQRPath != "" {
|
||||||
@@ -79,7 +93,7 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
part++
|
part++
|
||||||
}
|
}
|
||||||
|
|
||||||
if *pf.json {
|
if *pf.json && !qrToStdout {
|
||||||
b, _ := json.Marshal(jsonCerts)
|
b, _ := json.Marshal(jsonCerts)
|
||||||
_, _ = out.Write(b)
|
_, _ = out.Write(b)
|
||||||
_, _ = out.Write([]byte("\n"))
|
_, _ = out.Write([]byte("\n"))
|
||||||
@@ -91,7 +105,7 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return fmt.Errorf("error while generating qr code: %s", err)
|
return fmt.Errorf("error while generating qr code: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*pf.outQRPath, b, 0600)
|
err = writeOutput(*pf.outQRPath, b, 0600, out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-qr: %s", err)
|
return fmt.Errorf("error while writing out-qr: %s", err)
|
||||||
}
|
}
|
||||||
@@ -107,6 +121,7 @@ func printSummary() string {
|
|||||||
func printHelp(out io.Writer) {
|
func printHelp(out io.Writer) {
|
||||||
pf := newPrintFlags()
|
pf := newPrintFlags()
|
||||||
out.Write([]byte("Usage of " + os.Args[0] + " " + printSummary() + "\n"))
|
out.Write([]byte("Usage of " + os.Args[0] + " " + printSummary() + "\n"))
|
||||||
|
out.Write([]byte(stdioHelpText))
|
||||||
pf.set.SetOutput(out)
|
pf.set.SetOutput(out)
|
||||||
pf.set.PrintDefaults()
|
pf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ func Test_printHelp(t *testing.T) {
|
|||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
"Usage of "+os.Args[0]+" print <flags>: prints details about a certificate\n"+
|
"Usage of "+os.Args[0]+" print <flags>: prints details about a certificate\n"+
|
||||||
|
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
|
||||||
" -json\n"+
|
" -json\n"+
|
||||||
" \tOptional: outputs certificates in json format\n"+
|
" \tOptional: outputs certificates in json format\n"+
|
||||||
" -out-qr string\n"+
|
" -out-qr string\n"+
|
||||||
@@ -178,6 +179,44 @@ func Test_printCert(t *testing.T) {
|
|||||||
ob.String(),
|
ob.String(),
|
||||||
)
|
)
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
|
// read cert from stdin
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
withStdin(t, bytes.NewReader(p))
|
||||||
|
err = printCert([]string{"-json", "-path", "-"}, ob, eb)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(
|
||||||
|
t,
|
||||||
|
`[{"details":{"curve":"CURVE25519","groups":["hi"],"isCa":false,"issuer":"`+c.Issuer()+`","name":"test","networks":["10.0.0.123/8"],"notAfter":"0001-01-01T00:00:00Z","notBefore":"0001-01-01T00:00:00Z","publicKey":"`+pk+`","unsafeNetworks":[]},"fingerprint":"`+fp+`","signature":"`+sig+`","version":1}]
|
||||||
|
`,
|
||||||
|
ob.String(),
|
||||||
|
)
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
|
// -out-qr - sends only the PNG to stdout, suppressing the cert dump
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
withStdin(t, bytes.NewReader(p))
|
||||||
|
err = printCert([]string{"-path", "-", "-out-qr", "-"}, ob, eb)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
stdout := ob.Bytes()
|
||||||
|
require.NotEmpty(t, stdout)
|
||||||
|
// PNG magic, no PEM/JSON noise prepended
|
||||||
|
assert.Equal(t, []byte{0x89, 'P', 'N', 'G', 0x0d, 0x0a, 0x1a, 0x0a}, stdout[:8])
|
||||||
|
assert.NotContains(t, string(stdout), "NebulaCertificate")
|
||||||
|
assert.NotContains(t, string(stdout), `"details"`)
|
||||||
|
|
||||||
|
// json + out-qr - still suppresses json
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
withStdin(t, bytes.NewReader(p))
|
||||||
|
err = printCert([]string{"-json", "-path", "-", "-out-qr", "-"}, ob, eb)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
assert.Equal(t, []byte{0x89, 'P', 'N', 'G'}, ob.Bytes()[:4])
|
||||||
|
assert.NotContains(t, ob.String(), `"details"`)
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewTestCaCert will generate a CA cert
|
// NewTestCaCert will generate a CA cert
|
||||||
|
|||||||
+42
-20
@@ -85,6 +85,9 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
if !isP11 && *sf.inPubPath != "" && *sf.outKeyPath != "" {
|
if !isP11 && *sf.inPubPath != "" && *sf.outKeyPath != "" {
|
||||||
return newHelpErrorf("cannot set both -in-pub and -out-key")
|
return newHelpErrorf("cannot set both -in-pub and -out-key")
|
||||||
}
|
}
|
||||||
|
if isP11 && *sf.outKeyPath != "" {
|
||||||
|
return newHelpErrorf("cannot set -out-key with -pkcs11")
|
||||||
|
}
|
||||||
|
|
||||||
var v4Networks []netip.Prefix
|
var v4Networks []netip.Prefix
|
||||||
var v6Networks []netip.Prefix
|
var v6Networks []netip.Prefix
|
||||||
@@ -102,13 +105,35 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
return newHelpErrorf("-version must be either %v or %v", cert.Version1, cert.Version2)
|
return newHelpErrorf("-version must be either %v or %v", cert.Version1, cert.Version2)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if *sf.outKeyPath == "" {
|
||||||
|
*sf.outKeyPath = *sf.name + ".key"
|
||||||
|
}
|
||||||
|
if *sf.outCertPath == "" {
|
||||||
|
*sf.outCertPath = *sf.name + ".crt"
|
||||||
|
}
|
||||||
|
|
||||||
|
var claims ioClaims
|
||||||
|
if err := reserveInputs(&claims,
|
||||||
|
"ca-key", *sf.caKeyPath,
|
||||||
|
"ca-crt", *sf.caCertPath,
|
||||||
|
"in-pub", *sf.inPubPath,
|
||||||
|
); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := reserveOutputs(&claims,
|
||||||
|
"out-key", *sf.outKeyPath,
|
||||||
|
"out-crt", *sf.outCertPath,
|
||||||
|
"out-qr", *sf.outQRPath,
|
||||||
|
); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
var curve cert.Curve
|
var curve cert.Curve
|
||||||
var caKey []byte
|
var caKey []byte
|
||||||
|
|
||||||
if !isP11 {
|
if !isP11 {
|
||||||
var rawCAKey []byte
|
var rawCAKey []byte
|
||||||
rawCAKey, err := os.ReadFile(*sf.caKeyPath)
|
rawCAKey, err = readInput("ca-key", *sf.caKeyPath, &claims)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while reading ca-key: %s", err)
|
return fmt.Errorf("error while reading ca-key: %s", err)
|
||||||
}
|
}
|
||||||
@@ -121,7 +146,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
if len(passphrase) == 0 {
|
if len(passphrase) == 0 {
|
||||||
// ask for a passphrase until we get one
|
// ask for a passphrase until we get one
|
||||||
for i := 0; i < 5; i++ {
|
for i := 0; i < 5; i++ {
|
||||||
out.Write([]byte("Enter passphrase: "))
|
errOut.Write([]byte("Enter passphrase: "))
|
||||||
passphrase, err = pr.ReadPassword()
|
passphrase, err = pr.ReadPassword()
|
||||||
|
|
||||||
if errors.Is(err, ErrNoTerminal) {
|
if errors.Is(err, ErrNoTerminal) {
|
||||||
@@ -147,7 +172,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
rawCACert, err := os.ReadFile(*sf.caCertPath)
|
rawCACert, err := readInput("ca-crt", *sf.caCertPath, &claims)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while reading ca-crt: %s", err)
|
return fmt.Errorf("error while reading ca-crt: %s", err)
|
||||||
}
|
}
|
||||||
@@ -245,7 +270,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
|
|
||||||
if *sf.inPubPath != "" {
|
if *sf.inPubPath != "" {
|
||||||
var pubCurve cert.Curve
|
var pubCurve cert.Curve
|
||||||
rawPub, err := os.ReadFile(*sf.inPubPath)
|
rawPub, err := readInput("in-pub", *sf.inPubPath, &claims)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while reading in-pub: %s", err)
|
return fmt.Errorf("error while reading in-pub: %s", err)
|
||||||
}
|
}
|
||||||
@@ -266,16 +291,10 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
pub, rawPriv = newKeypair(curve)
|
pub, rawPriv = newKeypair(curve)
|
||||||
}
|
}
|
||||||
|
|
||||||
if *sf.outKeyPath == "" {
|
if !isStdio(*sf.outCertPath) {
|
||||||
*sf.outKeyPath = *sf.name + ".key"
|
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
||||||
}
|
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
||||||
|
}
|
||||||
if *sf.outCertPath == "" {
|
|
||||||
*sf.outCertPath = *sf.name + ".crt"
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
|
||||||
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var crts []cert.Certificate
|
var crts []cert.Certificate
|
||||||
@@ -360,11 +379,13 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !isP11 && *sf.inPubPath == "" {
|
if !isP11 && *sf.inPubPath == "" {
|
||||||
if _, err := os.Stat(*sf.outKeyPath); err == nil {
|
if !isStdio(*sf.outKeyPath) {
|
||||||
return fmt.Errorf("refusing to overwrite existing key: %s", *sf.outKeyPath)
|
if _, err := os.Stat(*sf.outKeyPath); err == nil {
|
||||||
|
return fmt.Errorf("refusing to overwrite existing key: %s", *sf.outKeyPath)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*sf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600)
|
err = writeOutput(*sf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600, out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-key: %s", err)
|
return fmt.Errorf("error while writing out-key: %s", err)
|
||||||
}
|
}
|
||||||
@@ -379,7 +400,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
b = append(b, sb...)
|
b = append(b, sb...)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*sf.outCertPath, b, 0600)
|
err = writeOutput(*sf.outCertPath, b, 0600, out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-crt: %s", err)
|
return fmt.Errorf("error while writing out-crt: %s", err)
|
||||||
}
|
}
|
||||||
@@ -390,7 +411,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
return fmt.Errorf("error while generating qr code: %s", err)
|
return fmt.Errorf("error while generating qr code: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*sf.outQRPath, b, 0600)
|
err = writeOutput(*sf.outQRPath, b, 0600, out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-qr: %s", err)
|
return fmt.Errorf("error while writing out-qr: %s", err)
|
||||||
}
|
}
|
||||||
@@ -440,6 +461,7 @@ func signSummary() string {
|
|||||||
func signHelp(out io.Writer) {
|
func signHelp(out io.Writer) {
|
||||||
sf := newSignFlags()
|
sf := newSignFlags()
|
||||||
out.Write([]byte("Usage of " + os.Args[0] + " " + signSummary() + "\n"))
|
out.Write([]byte("Usage of " + os.Args[0] + " " + signSummary() + "\n"))
|
||||||
|
out.Write([]byte(stdioHelpText))
|
||||||
sf.set.SetOutput(out)
|
sf.set.SetOutput(out)
|
||||||
sf.set.PrintDefaults()
|
sf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ func Test_signHelp(t *testing.T) {
|
|||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
"Usage of "+os.Args[0]+" sign <flags>: create and sign a certificate\n"+
|
"Usage of "+os.Args[0]+" sign <flags>: create and sign a certificate\n"+
|
||||||
|
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
|
||||||
" -ca-crt string\n"+
|
" -ca-crt string\n"+
|
||||||
" \tOptional: path to the signing CA cert (default \"ca.crt\")\n"+
|
" \tOptional: path to the signing CA cert (default \"ca.crt\")\n"+
|
||||||
" -ca-key string\n"+
|
" -ca-key string\n"+
|
||||||
@@ -376,15 +377,18 @@ func Test_signCert(t *testing.T) {
|
|||||||
// test with the proper password
|
// test with the proper password
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.NoError(t, signCert(args, ob, eb, testpw))
|
require.NoError(t, signCert(args, ob, eb, testpw))
|
||||||
assert.Equal(t, "Enter passphrase: ", ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "Enter passphrase: ", eb.String())
|
||||||
|
|
||||||
// test with the proper password in the environment
|
// test with the proper password in the environment
|
||||||
os.Remove(crtF.Name())
|
os.Remove(crtF.Name())
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
os.Setenv("NEBULA_CA_PASSPHRASE", string(passphrase))
|
os.Setenv("NEBULA_CA_PASSPHRASE", string(passphrase))
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
require.NoError(t, signCert(args, ob, eb, testpw))
|
require.NoError(t, signCert(args, ob, eb, testpw))
|
||||||
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
os.Setenv("NEBULA_CA_PASSPHRASE", "")
|
os.Setenv("NEBULA_CA_PASSPHRASE", "")
|
||||||
|
|
||||||
@@ -395,8 +399,8 @@ func Test_signCert(t *testing.T) {
|
|||||||
testpw.password = []byte("invalid password")
|
testpw.password = []byte("invalid password")
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.Error(t, signCert(args, ob, eb, testpw))
|
require.Error(t, signCert(args, ob, eb, testpw))
|
||||||
assert.Equal(t, "Enter passphrase: ", ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "Enter passphrase: ", eb.String())
|
||||||
|
|
||||||
// test with the wrong password in environment
|
// test with the wrong password in environment
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
@@ -416,8 +420,8 @@ func Test_signCert(t *testing.T) {
|
|||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.Error(t, signCert(args, ob, eb, nopw))
|
require.Error(t, signCert(args, ob, eb, nopw))
|
||||||
// normally the user hitting enter on the prompt would add newlines between these
|
// normally the user hitting enter on the prompt would add newlines between these
|
||||||
assert.Equal(t, "Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: ", ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: ", eb.String())
|
||||||
|
|
||||||
// test an error condition
|
// test an error condition
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
@@ -425,6 +429,106 @@ func Test_signCert(t *testing.T) {
|
|||||||
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.Error(t, signCert(args, ob, eb, errpw))
|
require.Error(t, signCert(args, ob, eb, errpw))
|
||||||
assert.Equal(t, "Enter passphrase: ", ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "Enter passphrase: ", eb.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_signCert_stdio(t *testing.T) {
|
||||||
|
nopw := &StubPasswordReader{
|
||||||
|
password: []byte(""),
|
||||||
|
err: nil,
|
||||||
|
}
|
||||||
|
|
||||||
|
caPub, caPriv, _ := ed25519.GenerateKey(rand.Reader)
|
||||||
|
rawCAKey := cert.MarshalSigningPrivateKeyToPEM(cert.Curve_CURVE25519, caPriv)
|
||||||
|
|
||||||
|
ca, _ := NewTestCaCert("ca", caPub, caPriv, time.Now(), time.Now().Add(time.Minute*200), nil, nil, nil)
|
||||||
|
rawCACrt, _ := ca.MarshalPEM()
|
||||||
|
|
||||||
|
caCrtF, err := os.CreateTemp("", "sign-cert.crt")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer os.Remove(caCrtF.Name())
|
||||||
|
caCrtF.Write(rawCACrt)
|
||||||
|
|
||||||
|
caKeyF, err := os.CreateTemp("", "sign-cert.key")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer os.Remove(caKeyF.Name())
|
||||||
|
caKeyF.Write(rawCAKey)
|
||||||
|
|
||||||
|
keyF, err := os.CreateTemp("", "sign.key")
|
||||||
|
require.NoError(t, err)
|
||||||
|
os.Remove(keyF.Name())
|
||||||
|
defer os.Remove(keyF.Name())
|
||||||
|
|
||||||
|
// ca-key on stdin, cert to stdout
|
||||||
|
withStdin(t, bytes.NewReader(rawCAKey))
|
||||||
|
ob := &bytes.Buffer{}
|
||||||
|
eb := &bytes.Buffer{}
|
||||||
|
args := []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", "-", "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", "-", "-out-key", keyF.Name(), "-duration", "100m"}
|
||||||
|
require.NoError(t, signCert(args, ob, eb, nopw))
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
|
lCrt, _, err := cert.UnmarshalCertificateFromPEM(ob.Bytes())
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "stdin-test", lCrt.Name())
|
||||||
|
assert.True(t, lCrt.CheckSignature(caPub))
|
||||||
|
|
||||||
|
// two flags reading from stdin should error before any read attempt;
|
||||||
|
// otherwise an interactive shell would hang on io.ReadAll
|
||||||
|
stdinIn := bytes.NewReader(rawCAKey)
|
||||||
|
withStdin(t, stdinIn)
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
args = []string{"-version", "1", "-ca-crt", "-", "-ca-key", "-", "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
||||||
|
require.EqualError(t, signCert(args, ob, eb, nopw),
|
||||||
|
`-ca-key and -ca-crt both set to "-", only one input may read from stdin`)
|
||||||
|
assert.Equal(t, len(rawCAKey), stdinIn.Len(), "stdin should be untouched when conflict is caught up front")
|
||||||
|
|
||||||
|
// two flags writing to stdout should error before any output is written
|
||||||
|
// AND before stdin is consumed
|
||||||
|
stdinR := bytes.NewReader(rawCAKey)
|
||||||
|
withStdin(t, stdinR)
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", "-", "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", "-", "-out-key", "-", "-duration", "100m"}
|
||||||
|
require.EqualError(t, signCert(args, ob, eb, nopw),
|
||||||
|
`-out-key and -out-crt both set to "-", only one output may write to stdout`)
|
||||||
|
assert.Empty(t, ob.String())
|
||||||
|
// stdin should be untouched because the conflict was caught up front
|
||||||
|
assert.Equal(t, len(rawCAKey), stdinR.Len())
|
||||||
|
|
||||||
|
// out-key on stdout, cert on disk
|
||||||
|
keyF2, err := os.CreateTemp("", "sign.key")
|
||||||
|
require.NoError(t, err)
|
||||||
|
os.Remove(keyF2.Name())
|
||||||
|
defer os.Remove(keyF2.Name())
|
||||||
|
crtF, err := os.CreateTemp("", "sign.crt")
|
||||||
|
require.NoError(t, err)
|
||||||
|
os.Remove(crtF.Name())
|
||||||
|
defer os.Remove(crtF.Name())
|
||||||
|
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", "-", "-duration", "100m"}
|
||||||
|
require.NoError(t, signCert(args, ob, eb, nopw))
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
_, _, curve, err := cert.UnmarshalPrivateKeyFromPEM(ob.Bytes())
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
||||||
|
|
||||||
|
// in-pub on stdin (caller already has a keypair, only the cert is generated)
|
||||||
|
inPub, _ := x25519Keypair()
|
||||||
|
rawInPub := cert.MarshalPublicKeyToPEM(cert.Curve_CURVE25519, inPub)
|
||||||
|
|
||||||
|
withStdin(t, bytes.NewReader(rawInPub))
|
||||||
|
os.Remove(crtF.Name())
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "in-pub-test", "-ip", "1.1.1.1/24", "-in-pub", "-", "-out-crt", "-", "-duration", "100m"}
|
||||||
|
require.NoError(t, signCert(args, ob, eb, nopw))
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
stdinCrt, _, err := cert.UnmarshalCertificateFromPEM(ob.Bytes())
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "in-pub-test", stdinCrt.Name())
|
||||||
|
assert.Equal(t, inPub, stdinCrt.PublicKey())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,117 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
)
|
||||||
|
|
||||||
|
// stdioPath is the special path value that selects stdin (for inputs) or
|
||||||
|
// stdout (for outputs) instead of a file on disk.
|
||||||
|
const stdioPath = "-"
|
||||||
|
|
||||||
|
// stdioHelpText is rendered just under the Usage line of each subcommand
|
||||||
|
// help so the - convention is documented once instead of on every flag.
|
||||||
|
const stdioHelpText = " Pass \"-\" to any path flag to read from stdin or write to stdout.\n"
|
||||||
|
|
||||||
|
// stdinReader is the source used when an input flag is set to "-".
|
||||||
|
// It is a package level var so tests can swap in a deterministic reader.
|
||||||
|
// Tests that mutate stdinReader cannot run with t.Parallel().
|
||||||
|
var stdinReader io.Reader = os.Stdin
|
||||||
|
|
||||||
|
// ioClaims tracks which flags have claimed stdin and stdout during a single
|
||||||
|
// command invocation so we can refuse a second flag asking for the same
|
||||||
|
// stream.
|
||||||
|
type ioClaims struct {
|
||||||
|
in string
|
||||||
|
out string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ioClaims) claimIn(flagName string) error {
|
||||||
|
if c.in != "" && c.in != flagName {
|
||||||
|
return fmt.Errorf("-%s and -%s both set to %q, only one input may read from stdin", c.in, flagName, stdioPath)
|
||||||
|
}
|
||||||
|
c.in = flagName
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ioClaims) claimOut(flagName string) error {
|
||||||
|
if c.out != "" && c.out != flagName {
|
||||||
|
return fmt.Errorf("-%s and -%s both set to %q, only one output may write to stdout", c.out, flagName, stdioPath)
|
||||||
|
}
|
||||||
|
c.out = flagName
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// reserveInputs walks alternating (flagName, path) pairs and claims stdin
|
||||||
|
// for any path equal to stdioPath. It must be called before any input is
|
||||||
|
// read so a conflict can be reported immediately instead of blocking on
|
||||||
|
// io.ReadAll while waiting for input that will never arrive.
|
||||||
|
func reserveInputs(claims *ioClaims, pairs ...string) error {
|
||||||
|
return reserveStdio(claims, "reserveInputs", (*ioClaims).claimIn, pairs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// reserveOutputs walks alternating (flagName, path) pairs and claims stdout
|
||||||
|
// for any path equal to stdioPath. It must be called before any output is
|
||||||
|
// written so a conflict cannot leave one stream half written before the
|
||||||
|
// second flag fails.
|
||||||
|
func reserveOutputs(claims *ioClaims, pairs ...string) error {
|
||||||
|
return reserveStdio(claims, "reserveOutputs", (*ioClaims).claimOut, pairs)
|
||||||
|
}
|
||||||
|
|
||||||
|
func reserveStdio(claims *ioClaims, who string, claim func(*ioClaims, string) error, pairs []string) error {
|
||||||
|
if len(pairs)%2 != 0 {
|
||||||
|
panic(who + " requires alternating name, path pairs")
|
||||||
|
}
|
||||||
|
for i := 0; i < len(pairs); i += 2 {
|
||||||
|
name, path := pairs[i], pairs[i+1]
|
||||||
|
if path != stdioPath {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := claim(claims, name); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// readInput returns the bytes referenced by path, reading from stdin when
|
||||||
|
// path is stdioPath.
|
||||||
|
func readInput(flagName, path string, claims *ioClaims) ([]byte, error) {
|
||||||
|
if path == stdioPath {
|
||||||
|
if err := claims.claimIn(flagName); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return io.ReadAll(stdinReader)
|
||||||
|
}
|
||||||
|
return os.ReadFile(path)
|
||||||
|
}
|
||||||
|
|
||||||
|
// openInput returns a reader for path. When path is stdioPath the returned
|
||||||
|
// reader wraps stdin and Close is a no-op.
|
||||||
|
func openInput(flagName, path string, claims *ioClaims) (io.ReadCloser, error) {
|
||||||
|
if path == stdioPath {
|
||||||
|
if err := claims.claimIn(flagName); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return io.NopCloser(stdinReader), nil
|
||||||
|
}
|
||||||
|
return os.Open(path)
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeOutput writes data to path, or to stdout when path is stdioPath. perm
|
||||||
|
// is only used for file output. The caller must have already claimed stdout
|
||||||
|
// via reserveOutputs before invoking with stdioPath.
|
||||||
|
func writeOutput(path string, data []byte, perm os.FileMode, stdout io.Writer) error {
|
||||||
|
if path == stdioPath {
|
||||||
|
_, err := stdout.Write(data)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return os.WriteFile(path, data, perm)
|
||||||
|
}
|
||||||
|
|
||||||
|
// isStdio reports whether path is the stdio sentinel and so should skip
|
||||||
|
// existence checks like "refuse to overwrite".
|
||||||
|
func isStdio(path string) bool {
|
||||||
|
return path == stdioPath
|
||||||
|
}
|
||||||
@@ -0,0 +1,167 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// withStdin temporarily replaces stdinReader for the duration of t.
|
||||||
|
func withStdin(t *testing.T, r io.Reader) {
|
||||||
|
t.Helper()
|
||||||
|
prev := stdinReader
|
||||||
|
stdinReader = r
|
||||||
|
t.Cleanup(func() { stdinReader = prev })
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_readInput_stdin(t *testing.T) {
|
||||||
|
withStdin(t, bytes.NewBufferString("hello"))
|
||||||
|
var claims ioClaims
|
||||||
|
|
||||||
|
got, err := readInput("path", "-", &claims)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, []byte("hello"), got)
|
||||||
|
assert.Equal(t, "path", claims.in)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_readInput_file(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
p := filepath.Join(dir, "f")
|
||||||
|
require.NoError(t, os.WriteFile(p, []byte("file"), 0600))
|
||||||
|
var claims ioClaims
|
||||||
|
|
||||||
|
got, err := readInput("path", p, &claims)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, []byte("file"), got)
|
||||||
|
assert.Empty(t, claims.in)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_readInput_doubleStdinErrors(t *testing.T) {
|
||||||
|
withStdin(t, bytes.NewBufferString("hello"))
|
||||||
|
var claims ioClaims
|
||||||
|
|
||||||
|
_, err := readInput("ca-key", "-", &claims)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
_, err = readInput("ca-crt", "-", &claims)
|
||||||
|
require.EqualError(t, err, `-ca-key and -ca-crt both set to "-", only one input may read from stdin`)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_openInput_stdin(t *testing.T) {
|
||||||
|
withStdin(t, bytes.NewBufferString("hi"))
|
||||||
|
var claims ioClaims
|
||||||
|
|
||||||
|
r, err := openInput("ca", "-", &claims)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer r.Close()
|
||||||
|
b, err := io.ReadAll(r)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, []byte("hi"), b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_openInput_doubleStdinErrors(t *testing.T) {
|
||||||
|
withStdin(t, bytes.NewBufferString("hi"))
|
||||||
|
var claims ioClaims
|
||||||
|
|
||||||
|
r, err := openInput("ca", "-", &claims)
|
||||||
|
require.NoError(t, err)
|
||||||
|
r.Close()
|
||||||
|
|
||||||
|
_, err = openInput("crt", "-", &claims)
|
||||||
|
require.EqualError(t, err, `-ca and -crt both set to "-", only one input may read from stdin`)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_writeOutput_stdout(t *testing.T) {
|
||||||
|
out := &bytes.Buffer{}
|
||||||
|
|
||||||
|
err := writeOutput("-", []byte("payload"), 0600, out)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "payload", out.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_writeOutput_file(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
p := filepath.Join(dir, "f")
|
||||||
|
out := &bytes.Buffer{}
|
||||||
|
|
||||||
|
err := writeOutput(p, []byte("payload"), 0600, out)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Empty(t, out.String())
|
||||||
|
got, err := os.ReadFile(p)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, []byte("payload"), got)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_reserveOutputs_noConflict(t *testing.T) {
|
||||||
|
var claims ioClaims
|
||||||
|
require.NoError(t, reserveOutputs(&claims,
|
||||||
|
"out-key", "/tmp/key",
|
||||||
|
"out-crt", "-",
|
||||||
|
"out-qr", "",
|
||||||
|
))
|
||||||
|
assert.Equal(t, "out-crt", claims.out)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_reserveOutputs_conflict(t *testing.T) {
|
||||||
|
var claims ioClaims
|
||||||
|
err := reserveOutputs(&claims,
|
||||||
|
"out-key", "-",
|
||||||
|
"out-crt", "-",
|
||||||
|
)
|
||||||
|
require.EqualError(t, err, `-out-key and -out-crt both set to "-", only one output may write to stdout`)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_reserveOutputs_panicsOnOddPairs(t *testing.T) {
|
||||||
|
defer func() {
|
||||||
|
r := recover()
|
||||||
|
require.NotNil(t, r)
|
||||||
|
}()
|
||||||
|
var claims ioClaims
|
||||||
|
_ = reserveOutputs(&claims, "out-key")
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_reserveInputs_noConflict(t *testing.T) {
|
||||||
|
var claims ioClaims
|
||||||
|
require.NoError(t, reserveInputs(&claims,
|
||||||
|
"ca-key", "/tmp/ca.key",
|
||||||
|
"ca-crt", "-",
|
||||||
|
"in-pub", "",
|
||||||
|
))
|
||||||
|
assert.Equal(t, "ca-crt", claims.in)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_reserveInputs_conflict(t *testing.T) {
|
||||||
|
var claims ioClaims
|
||||||
|
err := reserveInputs(&claims,
|
||||||
|
"ca-key", "-",
|
||||||
|
"ca-crt", "-",
|
||||||
|
)
|
||||||
|
require.EqualError(t, err, `-ca-key and -ca-crt both set to "-", only one input may read from stdin`)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_claimIn_idempotent(t *testing.T) {
|
||||||
|
// pre-claim then a lazy re-claim of the same flag should be a no-op
|
||||||
|
var claims ioClaims
|
||||||
|
require.NoError(t, claims.claimIn("ca-key"))
|
||||||
|
require.NoError(t, claims.claimIn("ca-key"))
|
||||||
|
assert.Equal(t, "ca-key", claims.in)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_claimOut_idempotent(t *testing.T) {
|
||||||
|
var claims ioClaims
|
||||||
|
require.NoError(t, claims.claimOut("out-crt"))
|
||||||
|
require.NoError(t, claims.claimOut("out-crt"))
|
||||||
|
assert.Equal(t, "out-crt", claims.out)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_isStdio(t *testing.T) {
|
||||||
|
assert.True(t, isStdio("-"))
|
||||||
|
assert.False(t, isStdio(""))
|
||||||
|
assert.False(t, isStdio("./-"))
|
||||||
|
assert.False(t, isStdio("foo"))
|
||||||
|
}
|
||||||
@@ -39,18 +39,26 @@ func verify(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
caFile, err := os.Open(*vf.caPath)
|
var claims ioClaims
|
||||||
|
if err := reserveInputs(&claims,
|
||||||
|
"ca", *vf.caPath,
|
||||||
|
"crt", *vf.certPath,
|
||||||
|
); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
caReader, err := openInput("ca", *vf.caPath, &claims)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while reading ca: %w", err)
|
return fmt.Errorf("error while reading ca: %w", err)
|
||||||
}
|
}
|
||||||
defer caFile.Close()
|
defer caReader.Close()
|
||||||
|
|
||||||
caPool, err := cert.NewCAPoolFromPEMReader(caFile)
|
caPool, err := cert.NewCAPoolFromPEMReader(caReader)
|
||||||
if err != nil && !errors.Is(err, cert.ErrExpired) {
|
if err != nil && !errors.Is(err, cert.ErrExpired) {
|
||||||
return fmt.Errorf("error while adding ca cert to pool: %w", err)
|
return fmt.Errorf("error while adding ca cert to pool: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
rawCert, err := os.ReadFile(*vf.certPath)
|
rawCert, err := readInput("crt", *vf.certPath, &claims)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("unable to read crt: %w", err)
|
return fmt.Errorf("unable to read crt: %w", err)
|
||||||
}
|
}
|
||||||
@@ -85,6 +93,7 @@ func verifySummary() string {
|
|||||||
func verifyHelp(out io.Writer) {
|
func verifyHelp(out io.Writer) {
|
||||||
vf := newVerifyFlags()
|
vf := newVerifyFlags()
|
||||||
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + verifySummary() + "\n"))
|
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + verifySummary() + "\n"))
|
||||||
|
_, _ = out.Write([]byte(stdioHelpText))
|
||||||
vf.set.SetOutput(out)
|
vf.set.SetOutput(out)
|
||||||
vf.set.PrintDefaults()
|
vf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ func Test_verifyHelp(t *testing.T) {
|
|||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
"Usage of "+os.Args[0]+" verify <flags>: verifies a certificate isn't expired and was signed by a trusted authority.\n"+
|
"Usage of "+os.Args[0]+" verify <flags>: verifies a certificate isn't expired and was signed by a trusted authority.\n"+
|
||||||
|
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
|
||||||
" -ca string\n"+
|
" -ca string\n"+
|
||||||
" \tRequired: path to a file containing one or more ca certificates\n"+
|
" \tRequired: path to a file containing one or more ca certificates\n"+
|
||||||
" -crt string\n"+
|
" -crt string\n"+
|
||||||
@@ -122,3 +123,46 @@ func Test_verify(t *testing.T) {
|
|||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func Test_verify_stdio(t *testing.T) {
|
||||||
|
ob := &bytes.Buffer{}
|
||||||
|
eb := &bytes.Buffer{}
|
||||||
|
|
||||||
|
caPub, caPriv, _ := ed25519.GenerateKey(rand.Reader)
|
||||||
|
ca, _ := NewTestCaCert("test-ca", caPub, caPriv, time.Now().Add(time.Hour*-1), time.Now().Add(time.Hour*2), nil, nil, nil)
|
||||||
|
caPEM, _ := ca.MarshalPEM()
|
||||||
|
|
||||||
|
crt, _ := NewTestCert(ca, caPriv, "test-cert", time.Now().Add(time.Hour*-1), time.Now().Add(time.Hour), nil, nil, nil)
|
||||||
|
crtPEM, _ := crt.MarshalPEM()
|
||||||
|
|
||||||
|
caFile, err := os.CreateTemp("", "verify-ca")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer os.Remove(caFile.Name())
|
||||||
|
caFile.Write(caPEM)
|
||||||
|
|
||||||
|
// crt on stdin, ca on disk
|
||||||
|
withStdin(t, bytes.NewReader(crtPEM))
|
||||||
|
require.NoError(t, verify([]string{"-ca", caFile.Name(), "-crt", "-"}, ob, eb))
|
||||||
|
assert.Empty(t, ob.String())
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
|
// ca on stdin, crt on disk
|
||||||
|
certFile, err := os.CreateTemp("", "verify-cert")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer os.Remove(certFile.Name())
|
||||||
|
certFile.Write(crtPEM)
|
||||||
|
|
||||||
|
withStdin(t, bytes.NewReader(caPEM))
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
require.NoError(t, verify([]string{"-ca", "-", "-crt", certFile.Name()}, ob, eb))
|
||||||
|
assert.Empty(t, ob.String())
|
||||||
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
|
// both flags on stdin should error
|
||||||
|
withStdin(t, bytes.NewReader(caPEM))
|
||||||
|
ob.Reset()
|
||||||
|
eb.Reset()
|
||||||
|
require.EqualError(t, verify([]string{"-ca", "-", "-crt", "-"}, ob, eb),
|
||||||
|
`-ca and -crt both set to "-", only one input may read from stdin`)
|
||||||
|
}
|
||||||
|
|||||||
@@ -53,7 +53,12 @@ func main() {
|
|||||||
l := logging.NewLogger(os.Stdout)
|
l := logging.NewLogger(os.Stdout)
|
||||||
|
|
||||||
if *serviceFlag != "" {
|
if *serviceFlag != "" {
|
||||||
if err := doService(configPath, configTest, Build, serviceFlag); err != nil {
|
if *configTest {
|
||||||
|
fmt.Println("-test is not supported with -service, run the config test without -service")
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := doService(configPath, Build, serviceFlag); err != nil {
|
||||||
l.Error("Service command failed", "error", err)
|
l.Error("Service command failed", "error", err)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
@@ -61,9 +66,12 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
fmt.Println("-config flag must be set")
|
p, err := config.DefaultPath()
|
||||||
flag.Usage()
|
if err != nil {
|
||||||
os.Exit(1)
|
fmt.Println(err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
*configPath = p
|
||||||
}
|
}
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
@@ -90,15 +98,14 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !*configTest {
|
if !*configTest {
|
||||||
wait, err := ctrl.Start()
|
if err := ctrl.Start(); err != nil {
|
||||||
if err != nil {
|
|
||||||
util.LogWithContextIfNeeded("Error while running", err, l)
|
util.LogWithContextIfNeeded("Error while running", err, l)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
go ctrl.ShutdownBlock()
|
go ctrl.ShutdownBlock()
|
||||||
|
|
||||||
if err := wait(); err != nil {
|
if err := ctrl.Wait(); err != nil {
|
||||||
l.Error("Nebula stopped due to fatal error", "error", err)
|
l.Error("Nebula stopped due to fatal error", "error", err)
|
||||||
os.Exit(2)
|
os.Exit(2)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
|
||||||
|
|
||||||
"github.com/kardianos/service"
|
"github.com/kardianos/service"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
@@ -16,7 +15,6 @@ var logger service.Logger
|
|||||||
|
|
||||||
type program struct {
|
type program struct {
|
||||||
configPath *string
|
configPath *string
|
||||||
configTest *bool
|
|
||||||
build string
|
build string
|
||||||
control *nebula.Control
|
control *nebula.Control
|
||||||
}
|
}
|
||||||
@@ -42,39 +40,47 @@ func (p *program) Start(s service.Service) error {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
p.control, err = nebula.Main(c, *p.configTest, Build, l, nil)
|
p.control, err = nebula.Main(c, false, Build, l, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
p.control.Start()
|
if err := p.control.Start(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Nebula can stop itself on a fatal packet reader error, make sure to log it if it happens.
|
||||||
|
go func() {
|
||||||
|
if err := p.control.Wait(); err != nil {
|
||||||
|
logger.Error(fmt.Sprintf("Nebula stopped due to fatal error: %v", err))
|
||||||
|
os.Exit(2)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *program) Stop(s service.Service) error {
|
func (p *program) Stop(s service.Service) error {
|
||||||
logger.Info("Nebula service stopping.")
|
logger.Info("Nebula service stopping.")
|
||||||
|
if p.control == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
p.control.Stop()
|
p.control.Stop()
|
||||||
|
|
||||||
|
// block until nebula has fully drained before reporting stopped.
|
||||||
|
// error logging is handled by Start.
|
||||||
|
_ = p.control.Wait()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func fileExists(filename string) bool {
|
func doService(configPath *string, build string, serviceFlag *string) error {
|
||||||
_, err := os.Stat(filename)
|
|
||||||
if os.IsNotExist(err) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func doService(configPath *string, configTest *bool, build string, serviceFlag *string) error {
|
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
ex, err := os.Executable()
|
p, err := config.DefaultPath()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
*configPath = filepath.Dir(ex) + "/config.yaml"
|
*configPath = p
|
||||||
if !fileExists(*configPath) {
|
|
||||||
*configPath = filepath.Dir(ex) + "/config.yml"
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
svcConfig := &service.Config{
|
svcConfig := &service.Config{
|
||||||
@@ -86,7 +92,6 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag *
|
|||||||
|
|
||||||
prg := &program{
|
prg := &program{
|
||||||
configPath: configPath,
|
configPath: configPath,
|
||||||
configTest: configTest,
|
|
||||||
build: build,
|
build: build,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -118,8 +123,9 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag *
|
|||||||
switch *serviceFlag {
|
switch *serviceFlag {
|
||||||
case "run":
|
case "run":
|
||||||
if err := s.Run(); err != nil {
|
if err := s.Run(); err != nil {
|
||||||
// Route any errors to the system logger
|
// Route any errors to the system logger and report the failure
|
||||||
logger.Error(err)
|
logger.Error(err)
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
default:
|
default:
|
||||||
if err := service.Control(s, *serviceFlag); err != nil {
|
if err := service.Control(s, *serviceFlag); err != nil {
|
||||||
|
|||||||
+8
-6
@@ -50,9 +50,12 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
fmt.Println("-config flag must be set")
|
p, err := config.DefaultPath()
|
||||||
flag.Usage()
|
if err != nil {
|
||||||
os.Exit(1)
|
fmt.Println(err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
*configPath = p
|
||||||
}
|
}
|
||||||
|
|
||||||
l := logging.NewLogger(os.Stdout)
|
l := logging.NewLogger(os.Stdout)
|
||||||
@@ -81,8 +84,7 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !*configTest {
|
if !*configTest {
|
||||||
wait, err := ctrl.Start()
|
if err := ctrl.Start(); err != nil {
|
||||||
if err != nil {
|
|
||||||
util.LogWithContextIfNeeded("Error while running", err, l)
|
util.LogWithContextIfNeeded("Error while running", err, l)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
@@ -90,7 +92,7 @@ func main() {
|
|||||||
go ctrl.ShutdownBlock()
|
go ctrl.ShutdownBlock()
|
||||||
notifyReady(l)
|
notifyReady(l)
|
||||||
|
|
||||||
if err := wait(); err != nil {
|
if err := ctrl.Wait(); err != nil {
|
||||||
l.Error("Nebula stopped due to fatal error", "error", err)
|
l.Error("Nebula stopped due to fatal error", "error", err)
|
||||||
os.Exit(2)
|
os.Exit(2)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,29 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DefaultPath returns a path to a config file alongside the running executable, preferring config.yaml over config.yml.
|
||||||
|
// If neither file exists an error is returned that names both paths checked.
|
||||||
|
func DefaultPath() (string, error) {
|
||||||
|
ex, err := os.Executable()
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return defaultPathInDir(filepath.Dir(ex))
|
||||||
|
}
|
||||||
|
|
||||||
|
func defaultPathInDir(dir string) (string, error) {
|
||||||
|
yamlPath := filepath.Join(dir, "config.yaml")
|
||||||
|
if _, err := os.Stat(yamlPath); err == nil {
|
||||||
|
return yamlPath, nil
|
||||||
|
}
|
||||||
|
ymlPath := filepath.Join(dir, "config.yml")
|
||||||
|
if _, err := os.Stat(ymlPath); err == nil {
|
||||||
|
return ymlPath, nil
|
||||||
|
}
|
||||||
|
return "", fmt.Errorf("no default config found at %s or %s", yamlPath, ymlPath)
|
||||||
|
}
|
||||||
@@ -0,0 +1,67 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestDefaultPathInDir(t *testing.T) {
|
||||||
|
t.Run("prefers config.yaml when both exist", func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
want := filepath.Join(dir, "config.yaml")
|
||||||
|
other := filepath.Join(dir, "config.yml")
|
||||||
|
require.NoError(t, os.WriteFile(want, []byte("a: 1"), 0644))
|
||||||
|
require.NoError(t, os.WriteFile(other, []byte("a: 2"), 0644))
|
||||||
|
|
||||||
|
got, err := defaultPathInDir(dir)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, want, got)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("returns config.yaml when only it exists", func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
want := filepath.Join(dir, "config.yaml")
|
||||||
|
require.NoError(t, os.WriteFile(want, []byte("a: 1"), 0644))
|
||||||
|
|
||||||
|
got, err := defaultPathInDir(dir)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, want, got)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("falls back to config.yml when only it exists", func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
want := filepath.Join(dir, "config.yml")
|
||||||
|
require.NoError(t, os.WriteFile(want, []byte("a: 1"), 0644))
|
||||||
|
|
||||||
|
got, err := defaultPathInDir(dir)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, want, got)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("errors when neither exists and names both paths", func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
got, err := defaultPathInDir(dir)
|
||||||
|
assert.Empty(t, got)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), filepath.Join(dir, "config.yaml"))
|
||||||
|
assert.Contains(t, err.Error(), filepath.Join(dir, "config.yml"))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDefaultPath(t *testing.T) {
|
||||||
|
got, err := DefaultPath()
|
||||||
|
if err != nil {
|
||||||
|
ex, exErr := os.Executable()
|
||||||
|
require.NoError(t, exErr)
|
||||||
|
assert.Contains(t, err.Error(), filepath.Dir(ex))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ex, err := os.Executable()
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, filepath.Dir(ex), filepath.Dir(got))
|
||||||
|
assert.Contains(t, []string{"config.yaml", "config.yml"}, filepath.Base(got))
|
||||||
|
}
|
||||||
+14
-52
@@ -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)
|
||||||
@@ -140,14 +136,6 @@ func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time)
|
|||||||
return in, out
|
return in, out
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddTrafficWatch must be called for every new HostInfo.
|
|
||||||
// We will continue to monitor the HostInfo until the tunnel is dropped.
|
|
||||||
func (cm *connectionManager) AddTrafficWatch(h *HostInfo) {
|
|
||||||
if h.out.Swap(true) == false {
|
|
||||||
cm.trafficTimer.Add(h.localIndexId, cm.checkInterval)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (cm *connectionManager) Start(ctx context.Context) {
|
func (cm *connectionManager) Start(ctx context.Context) {
|
||||||
clockSource := time.NewTicker(cm.trafficTimer.t.tickDuration)
|
clockSource := time.NewTicker(cm.trafficTimer.t.tickDuration)
|
||||||
defer clockSource.Stop()
|
defer clockSource.Stop()
|
||||||
@@ -310,8 +298,8 @@ func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo
|
|||||||
} else {
|
} else {
|
||||||
cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu))
|
cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||||
cm.l.Info("send CreateRelayRequest",
|
cm.l.Info("send CreateRelayRequest",
|
||||||
"relayFrom", req.RelayFromAddr,
|
"relayFrom", relayFrom,
|
||||||
"relayTo", req.RelayToAddr,
|
"relayTo", relayTo,
|
||||||
"initiatorRelayIndex", req.InitiatorRelayIndex,
|
"initiatorRelayIndex", req.InitiatorRelayIndex,
|
||||||
"responderRelayIndex", req.ResponderRelayIndex,
|
"responderRelayIndex", req.ResponderRelayIndex,
|
||||||
"vpnAddrs", newhostinfo.vpnAddrs,
|
"vpnAddrs", newhostinfo.vpnAddrs,
|
||||||
@@ -369,7 +357,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 +388,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 +499,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
|
||||||
|
|||||||
+6
-5
@@ -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 = 8192
|
||||||
|
|
||||||
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)
|
||||||
|
|||||||
+51
-25
@@ -69,29 +69,29 @@ type ControlHostInfo struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Start actually runs nebula, this is a nonblocking call.
|
// Start actually runs nebula, this is a nonblocking call.
|
||||||
// The returned function blocks until nebula has fully stopped and returns the
|
// Use Wait to block until nebula has fully stopped and to learn whether a fatal reader error caused the shutdown.
|
||||||
// first fatal reader error (if any). A nil error means nebula shut down
|
func (c *Control) Start() error {
|
||||||
// gracefully; a non-nil error means a reader hit an unexpected failure that
|
|
||||||
// triggered the shutdown.
|
|
||||||
func (c *Control) Start() (func() error, error) {
|
|
||||||
c.stateLock.Lock()
|
c.stateLock.Lock()
|
||||||
defer c.stateLock.Unlock()
|
defer c.stateLock.Unlock()
|
||||||
switch c.state {
|
switch c.state {
|
||||||
case StateReady:
|
case StateReady:
|
||||||
//yay!
|
//yay!
|
||||||
case StateStopped, StateStopping:
|
case StateStopped, StateStopping:
|
||||||
return nil, ErrAlreadyStopped
|
return ErrAlreadyStopped
|
||||||
case StateStarted:
|
case StateStarted:
|
||||||
return nil, ErrAlreadyStarted
|
return ErrAlreadyStarted
|
||||||
default:
|
default:
|
||||||
return nil, ErrUnknownState
|
return ErrUnknownState
|
||||||
}
|
}
|
||||||
|
|
||||||
// Activate the interface
|
// Activate the interface
|
||||||
err := c.f.activate()
|
err := c.f.activate()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
// Cancel before Close so a caller returning from Wait always observes a dead Context
|
||||||
|
c.cancel()
|
||||||
|
_ = c.f.Close()
|
||||||
c.state = StateStopped
|
c.state = StateStopped
|
||||||
return nil, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Call all the delayed funcs that waited patiently for the interface to be created.
|
// Call all the delayed funcs that waited patiently for the interface to be created.
|
||||||
@@ -114,13 +114,9 @@ func (c *Control) Start() (func() error, error) {
|
|||||||
c.f.triggerShutdown = c.Stop
|
c.f.triggerShutdown = c.Stop
|
||||||
|
|
||||||
// Start reading packets.
|
// Start reading packets.
|
||||||
out, err := c.f.run()
|
c.f.run()
|
||||||
if err != nil {
|
|
||||||
c.state = StateStopped
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
c.state = StateStarted
|
c.state = StateStarted
|
||||||
return out, nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) State() RunState {
|
func (c *Control) State() RunState {
|
||||||
@@ -133,10 +129,26 @@ func (c *Control) Context() context.Context {
|
|||||||
return c.ctx
|
return c.ctx
|
||||||
}
|
}
|
||||||
|
|
||||||
// Stop is a non-blocking call that signals nebula to close all tunnels and shut down
|
// Stop tears nebula down, closing all tunnels and releasing everything it holds.
|
||||||
|
// Use Wait to block until the shutdown has completed.
|
||||||
|
// A Control that has been stopped cannot be started again, Start will return ErrAlreadyStopped.
|
||||||
func (c *Control) Stop() {
|
func (c *Control) Stop() {
|
||||||
c.stateLock.Lock()
|
c.stateLock.Lock()
|
||||||
if c.state != StateStarted {
|
switch c.state {
|
||||||
|
case StateStarted:
|
||||||
|
// Fall through to the full teardown below
|
||||||
|
|
||||||
|
case StateReady:
|
||||||
|
// Never started
|
||||||
|
c.cancel()
|
||||||
|
c.state = StateStopped
|
||||||
|
if err := c.f.Close(); err != nil {
|
||||||
|
c.l.Error("Close interface failed", "error", err)
|
||||||
|
}
|
||||||
|
c.stateLock.Unlock()
|
||||||
|
return
|
||||||
|
|
||||||
|
default:
|
||||||
c.stateLock.Unlock()
|
c.stateLock.Unlock()
|
||||||
// We are stopping or stopped already
|
// We are stopping or stopped already
|
||||||
return
|
return
|
||||||
@@ -145,19 +157,26 @@ func (c *Control) Stop() {
|
|||||||
c.state = StateStopping
|
c.state = StateStopping
|
||||||
c.stateLock.Unlock()
|
c.stateLock.Unlock()
|
||||||
|
|
||||||
// Stop the handshakeManager (and other services), to prevent new tunnels from
|
// Closing tunnels can be slow with a large hostmap, don't hold the lock for it
|
||||||
// being created while we're shutting them all down.
|
|
||||||
c.cancel()
|
c.cancel()
|
||||||
|
|
||||||
c.CloseAllTunnels(false)
|
c.CloseAllTunnels(false)
|
||||||
|
|
||||||
|
c.stateLock.Lock()
|
||||||
|
c.state = StateStopped
|
||||||
if err := c.f.Close(); err != nil {
|
if err := c.f.Close(); err != nil {
|
||||||
c.l.Error("Close interface failed", "error", err)
|
c.l.Error("Close interface failed", "error", err)
|
||||||
}
|
}
|
||||||
c.stateLock.Lock()
|
|
||||||
c.state = StateStopped
|
|
||||||
c.stateLock.Unlock()
|
c.stateLock.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Wait blocks until nebula has fully stopped, either via Stop or an internal fatal error,
|
||||||
|
// and returns the first fatal packet reader error if there was one.
|
||||||
|
// It is safe to call from multiple goroutines and at any point in the lifecycle,
|
||||||
|
// but a Wait on a Control that is never started and never stopped will block forever.
|
||||||
|
func (c *Control) Wait() error {
|
||||||
|
return c.f.wait()
|
||||||
|
}
|
||||||
|
|
||||||
// ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled
|
// ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled
|
||||||
func (c *Control) ShutdownBlock() {
|
func (c *Control) ShutdownBlock() {
|
||||||
sigChan := make(chan os.Signal, 1)
|
sigChan := make(chan os.Signal, 1)
|
||||||
@@ -170,8 +189,15 @@ func (c *Control) ShutdownBlock() {
|
|||||||
c.Stop()
|
c.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change
|
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change.
|
||||||
func (c *Control) RebindUDPServer() {
|
func (c *Control) RebindUDPServer() {
|
||||||
|
c.stateLock.Lock()
|
||||||
|
defer c.stateLock.Unlock()
|
||||||
|
|
||||||
|
if c.state != StateStarted {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
_ = c.f.outside.Rebind()
|
_ = c.f.outside.Rebind()
|
||||||
|
|
||||||
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
|
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
|
||||||
@@ -305,7 +331,7 @@ func (c *Control) CloseAllTunnels(excludeLighthouses bool) (closed int) {
|
|||||||
|
|
||||||
c.l.Debug("Sending close tunnel message",
|
c.l.Debug("Sending close tunnel message",
|
||||||
"vpnAddrs", h.vpnAddrs,
|
"vpnAddrs", h.vpnAddrs,
|
||||||
"udpAddr", h.remote,
|
"udpAddr", h.GetRemote(),
|
||||||
)
|
)
|
||||||
closed++
|
closed++
|
||||||
}
|
}
|
||||||
@@ -350,7 +376,7 @@ func copyHostInfo(h *HostInfo, preferredRanges []netip.Prefix) ControlHostInfo {
|
|||||||
RemoteAddrs: h.remotes.CopyAddrs(preferredRanges),
|
RemoteAddrs: h.remotes.CopyAddrs(preferredRanges),
|
||||||
CurrentRelaysToMe: h.relayState.CopyRelayIps(),
|
CurrentRelaysToMe: h.relayState.CopyRelayIps(),
|
||||||
CurrentRelaysThroughMe: h.relayState.CopyRelayForIps(),
|
CurrentRelaysThroughMe: h.relayState.CopyRelayForIps(),
|
||||||
CurrentRemote: h.remote,
|
CurrentRemote: h.GetRemote(),
|
||||||
}
|
}
|
||||||
|
|
||||||
for i, a := range h.vpnAddrs {
|
for i, a := range h.vpnAddrs {
|
||||||
|
|||||||
@@ -0,0 +1,309 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/netip"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gaissmai/bart"
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
|
"github.com/slackhq/nebula/routing"
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
type fakeDevice struct {
|
||||||
|
closeOnce sync.Once
|
||||||
|
closedCh chan struct{}
|
||||||
|
closed bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFakeDevice() *fakeDevice {
|
||||||
|
return &fakeDevice{closedCh: make(chan struct{})}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read blocks until Close like a real tun with no traffic, then reports EOF
|
||||||
|
// the same way a closed device does
|
||||||
|
func (d *fakeDevice) Read() ([]tio.Packet, error) {
|
||||||
|
<-d.closedCh
|
||||||
|
return nil, io.EOF
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *fakeDevice) Write(p []byte) (int, error) { return len(p), nil }
|
||||||
|
|
||||||
|
func (d *fakeDevice) Close() error {
|
||||||
|
d.closeOnce.Do(func() {
|
||||||
|
d.closed = true
|
||||||
|
close(d.closedCh)
|
||||||
|
})
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *fakeDevice) Activate() error { return nil }
|
||||||
|
func (d *fakeDevice) Networks() []netip.Prefix { return nil }
|
||||||
|
func (d *fakeDevice) Name() string { return "fake" }
|
||||||
|
func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil }
|
||||||
|
|
||||||
|
func (d *fakeDevice) Queues(int) ([]tio.Queue, error) { return []tio.Queue{d}, nil }
|
||||||
|
|
||||||
|
// newReadyControl hand-builds the minimum Control that Main would have
|
||||||
|
// produced right before Start, including the construction token NewInterface
|
||||||
|
// takes so waiters block until Close releases the resources
|
||||||
|
func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
dev := newFakeDevice()
|
||||||
|
conn := &fakeConn{}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
|
||||||
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
||||||
|
nt := new(bart.Lite)
|
||||||
|
nt.Insert(myVpnNet)
|
||||||
|
cs := &CertState{
|
||||||
|
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||||
|
myVpnNetworksTable: nt,
|
||||||
|
}
|
||||||
|
lh, err := NewLightHouseFromConfig(ctx, l, config.NewC(l), cs, nil, nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
f := &Interface{
|
||||||
|
ctx: ctx,
|
||||||
|
inside: dev,
|
||||||
|
outside: conn,
|
||||||
|
writers: []udp.Conn{conn},
|
||||||
|
batchers: make([]batch.RxBatcher, 1),
|
||||||
|
routines: 1,
|
||||||
|
hostMap: newHostMap(l),
|
||||||
|
lightHouse: lh,
|
||||||
|
l: l,
|
||||||
|
}
|
||||||
|
f.wg.Add(1)
|
||||||
|
|
||||||
|
return &Control{
|
||||||
|
state: StateReady,
|
||||||
|
f: f,
|
||||||
|
l: l,
|
||||||
|
ctx: ctx,
|
||||||
|
cancel: cancel,
|
||||||
|
}, dev, conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_StopBeforeStart(t *testing.T) {
|
||||||
|
c, dev, conn := newReadyControl(t)
|
||||||
|
|
||||||
|
// A Stop on a never started control must release everything Main acquired
|
||||||
|
c.Stop()
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||||
|
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||||
|
require.ErrorIs(t, c.ctx.Err(), context.Canceled, "the service context should have been cancelled")
|
||||||
|
|
||||||
|
// Wait must return promptly now that the resources are released
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
|
||||||
|
// A stopped control can never be started
|
||||||
|
err := c.Start()
|
||||||
|
require.ErrorIs(t, err, ErrAlreadyStopped)
|
||||||
|
|
||||||
|
// A second Stop is a harmless no-op
|
||||||
|
c.Stop()
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_WaitBlocksUntilStop(t *testing.T) {
|
||||||
|
c, _, _ := newReadyControl(t)
|
||||||
|
|
||||||
|
done := make(chan error, 1)
|
||||||
|
go func() { done <- c.Wait() }()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
t.Fatal("Wait returned before Stop")
|
||||||
|
case <-time.After(50 * time.Millisecond):
|
||||||
|
}
|
||||||
|
|
||||||
|
c.Stop()
|
||||||
|
select {
|
||||||
|
case err := <-done:
|
||||||
|
require.NoError(t, err)
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("Wait did not return after Stop")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type fakeConn struct {
|
||||||
|
closed bool
|
||||||
|
rebinds int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
|
||||||
|
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||||
|
func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil }
|
||||||
|
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
||||||
|
func (c *fakeConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
||||||
|
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
||||||
|
func (c *fakeConn) Close() error { c.closed = true; return nil }
|
||||||
|
|
||||||
|
type multiqueueDevice struct {
|
||||||
|
*fakeDevice
|
||||||
|
}
|
||||||
|
|
||||||
|
// Queues claims multiqueue support but fails to open the second queue,
|
||||||
|
// exercising the activation error path.
|
||||||
|
func (d *multiqueueDevice) Queues(n int) ([]tio.Queue, error) {
|
||||||
|
if n > 1 {
|
||||||
|
return nil, errors.New("second queue failed to open")
|
||||||
|
}
|
||||||
|
return d.fakeDevice.Queues(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
||||||
|
dev := &multiqueueDevice{fakeDevice: newFakeDevice()}
|
||||||
|
conn := &fakeConn{}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
f := &Interface{
|
||||||
|
ctx: ctx,
|
||||||
|
inside: dev,
|
||||||
|
outside: conn,
|
||||||
|
writers: []udp.Conn{conn},
|
||||||
|
batchers: make([]batch.RxBatcher, 2),
|
||||||
|
routines: 2,
|
||||||
|
l: test.NewLogger(),
|
||||||
|
}
|
||||||
|
f.wg.Add(1)
|
||||||
|
|
||||||
|
c := &Control{
|
||||||
|
state: StateReady,
|
||||||
|
f: f,
|
||||||
|
l: test.NewLogger(),
|
||||||
|
ctx: ctx,
|
||||||
|
cancel: cancel,
|
||||||
|
}
|
||||||
|
|
||||||
|
// The second reader fails to open, everything must be released
|
||||||
|
err := c.Start()
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||||
|
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||||
|
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
|
||||||
|
|
||||||
|
// And Wait must not hang on the construction token
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInterface_CloseIsIdempotent(t *testing.T) {
|
||||||
|
dev := newFakeDevice()
|
||||||
|
f := &Interface{
|
||||||
|
inside: dev,
|
||||||
|
l: test.NewLogger(),
|
||||||
|
}
|
||||||
|
f.wg.Add(1)
|
||||||
|
|
||||||
|
require.NoError(t, f.Close())
|
||||||
|
assert.True(t, dev.closed)
|
||||||
|
|
||||||
|
// A second Close must not double release the wg token or the device
|
||||||
|
require.NoError(t, f.Close())
|
||||||
|
require.NoError(t, f.wait())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_FatalErrorReportsThroughWait(t *testing.T) {
|
||||||
|
c, dev, conn := newReadyControl(t)
|
||||||
|
|
||||||
|
// Mirror what Start wires up, without needing real packet readers
|
||||||
|
c.f.triggerShutdown = c.Stop
|
||||||
|
c.state = StateStarted
|
||||||
|
|
||||||
|
boom := errors.New("boom")
|
||||||
|
c.f.onFatal(boom)
|
||||||
|
|
||||||
|
require.ErrorIs(t, c.Wait(), boom)
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
assert.True(t, dev.closed)
|
||||||
|
assert.True(t, conn.closed)
|
||||||
|
|
||||||
|
// A second fatal error must not fire the shutdown again or replace the first
|
||||||
|
c.f.onFatal(errors.New("later"))
|
||||||
|
require.ErrorIs(t, c.Wait(), boom)
|
||||||
|
|
||||||
|
// Wait stays factual, a Stop after the death does not mask the error
|
||||||
|
c.Stop()
|
||||||
|
require.ErrorIs(t, c.Wait(), boom)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_ConcurrentStopAndStart(t *testing.T) {
|
||||||
|
c, _, _ := newReadyControl(t)
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
wg.Go(func() { c.Stop() })
|
||||||
|
}
|
||||||
|
wg.Go(func() { _ = c.Start() })
|
||||||
|
wg.Go(func() {
|
||||||
|
_ = c.Wait()
|
||||||
|
// A returned Wait must always observe the final state, no matter how
|
||||||
|
// the race resolved
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
})
|
||||||
|
wg.Wait()
|
||||||
|
|
||||||
|
// However the race resolves, the control must end fully stopped with no
|
||||||
|
// panic and Wait must observe the final state
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
err := c.Start()
|
||||||
|
require.ErrorIs(t, err, ErrAlreadyStopped)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_StartStopLifecycle(t *testing.T) {
|
||||||
|
c, dev, conn := newReadyControl(t)
|
||||||
|
|
||||||
|
err := c.Start()
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, StateStarted, c.State())
|
||||||
|
err = c.Start()
|
||||||
|
require.ErrorIs(t, err, ErrAlreadyStarted)
|
||||||
|
|
||||||
|
// Stop must unpark the reader blocked in the device and release everything
|
||||||
|
c.Stop()
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||||
|
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||||
|
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
|
||||||
|
|
||||||
|
// The reader drained off a closed device, that is not a fatal error
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
err = c.Start()
|
||||||
|
require.ErrorIs(t, err, ErrAlreadyStopped)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_RebindIsGatedByState(t *testing.T) {
|
||||||
|
c, _, conn := newReadyControl(t)
|
||||||
|
|
||||||
|
// A rebind before Start reaches nothing, the interface is not up
|
||||||
|
c.RebindUDPServer()
|
||||||
|
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
|
||||||
|
|
||||||
|
err := c.Start()
|
||||||
|
require.NoError(t, err)
|
||||||
|
c.RebindUDPServer()
|
||||||
|
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
|
||||||
|
|
||||||
|
// A rebind racing a completed stop must not touch the closed conn
|
||||||
|
c.Stop()
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
c.RebindUDPServer()
|
||||||
|
assert.Equal(t, 1, conn.rebinds, "rebind after stop must be a no-op")
|
||||||
|
}
|
||||||
+161
-6
@@ -1,6 +1,8 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"reflect"
|
"reflect"
|
||||||
@@ -9,6 +11,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
||||||
@@ -42,8 +45,7 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
|||||||
assert.True(t, ok)
|
assert.True(t, ok)
|
||||||
|
|
||||||
crt := &dummyCert{}
|
crt := &dummyCert{}
|
||||||
hm.unlockedAddHostInfo(&HostInfo{
|
hi := &HostInfo{
|
||||||
remote: remote1,
|
|
||||||
remotes: remotes,
|
remotes: remotes,
|
||||||
ConnectionState: &ConnectionState{
|
ConnectionState: &ConnectionState{
|
||||||
peerCert: &cert.CachedCertificate{Certificate: crt},
|
peerCert: &cert.CachedCertificate{Certificate: crt},
|
||||||
@@ -56,13 +58,14 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
|||||||
relayForByAddr: map[netip.Addr]*Relay{},
|
relayForByAddr: map[netip.Addr]*Relay{},
|
||||||
relayForByIdx: map[uint32]*Relay{},
|
relayForByIdx: map[uint32]*Relay{},
|
||||||
},
|
},
|
||||||
}, &Interface{})
|
}
|
||||||
|
hi.remote.Store(&remote1)
|
||||||
|
hm.unlockedAddHostInfo(hi, &Interface{})
|
||||||
|
|
||||||
vpnIp2, ok := netip.AddrFromSlice(ipNet2.IP)
|
vpnIp2, ok := netip.AddrFromSlice(ipNet2.IP)
|
||||||
assert.True(t, ok)
|
assert.True(t, ok)
|
||||||
|
|
||||||
hm.unlockedAddHostInfo(&HostInfo{
|
hi2 := &HostInfo{
|
||||||
remote: remote1,
|
|
||||||
remotes: remotes,
|
remotes: remotes,
|
||||||
ConnectionState: &ConnectionState{
|
ConnectionState: &ConnectionState{
|
||||||
peerCert: nil,
|
peerCert: nil,
|
||||||
@@ -75,7 +78,9 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
|||||||
relayForByAddr: map[netip.Addr]*Relay{},
|
relayForByAddr: map[netip.Addr]*Relay{},
|
||||||
relayForByIdx: map[uint32]*Relay{},
|
relayForByIdx: map[uint32]*Relay{},
|
||||||
},
|
},
|
||||||
}, &Interface{})
|
}
|
||||||
|
hi2.remote.Store(&remote1)
|
||||||
|
hm.unlockedAddHostInfo(hi2, &Interface{})
|
||||||
|
|
||||||
c := Control{
|
c := Control{
|
||||||
state: StateReady,
|
state: StateReady,
|
||||||
@@ -119,3 +124,153 @@ func assertFields(t *testing.T, expected []string, actualStruct any) {
|
|||||||
|
|
||||||
assert.Equal(t, expected, fields)
|
assert.Equal(t, expected, fields)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// alwaysAllowV4/V6 are check funcs that accept every entry (including nil pointers),
|
||||||
|
// letting us inject a nil *V4AddrPort/*V6AddrPort into a RemoteList's reported cache
|
||||||
|
// the same way a malformed proto message off the wire could.
|
||||||
|
func alwaysAllowV4(netip.Addr, *V4AddrPort) bool { return true }
|
||||||
|
func alwaysAllowV6(netip.Addr, *V6AddrPort) bool { return true }
|
||||||
|
|
||||||
|
// TestGetRelays_SkipsNilRelayAddrs proves GetRelays tolerates nil entries in the
|
||||||
|
// RelayVpnAddrs proto slice (which protoAddrToNetAddr would nil-deref on) and still
|
||||||
|
// returns the valid relays, including the legacy OldRelayVpnAddrs.
|
||||||
|
func TestGetRelays_SkipsNilRelayAddrs(t *testing.T) {
|
||||||
|
good := netip.MustParseAddr("10.0.0.9")
|
||||||
|
|
||||||
|
d := &NebulaMetaDetails{
|
||||||
|
OldRelayVpnAddrs: []uint32{0x0a000001}, // 10.0.0.1
|
||||||
|
RelayVpnAddrs: []*Addr{
|
||||||
|
nil,
|
||||||
|
netAddrToProtoAddr(good),
|
||||||
|
nil,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
var relays []netip.Addr
|
||||||
|
require.NotPanics(t, func() { relays = d.GetRelays() })
|
||||||
|
|
||||||
|
assert.Equal(t, []netip.Addr{
|
||||||
|
netip.MustParseAddr("10.0.0.1"),
|
||||||
|
good,
|
||||||
|
}, relays)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestGetRelays_AllNil ensures an all-nil RelayVpnAddrs slice yields no relays and no panic.
|
||||||
|
func TestGetRelays_AllNil(t *testing.T) {
|
||||||
|
d := &NebulaMetaDetails{RelayVpnAddrs: []*Addr{nil, nil}}
|
||||||
|
var relays []netip.Addr
|
||||||
|
require.NotPanics(t, func() { relays = d.GetRelays() })
|
||||||
|
assert.Empty(t, relays)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRemoteList_CopyCache_SkipsNilReported proves CopyCache skips nil reported
|
||||||
|
// pointers (v4 and v6) instead of nil-dereferencing them in protoV*AddrPortToNetAddrPort.
|
||||||
|
func TestRemoteList_CopyCache_SkipsNilReported(t *testing.T) {
|
||||||
|
owner := netip.MustParseAddr("10.0.0.1")
|
||||||
|
rl := NewRemoteList([]netip.Addr{owner}, nil)
|
||||||
|
|
||||||
|
rl.unlockedSetV4(owner, owner, []*V4AddrPort{
|
||||||
|
nil,
|
||||||
|
newIp4AndPortFromString("1.2.3.4:5"),
|
||||||
|
nil,
|
||||||
|
}, alwaysAllowV4)
|
||||||
|
|
||||||
|
rl.unlockedSetV6(owner, owner, []*V6AddrPort{
|
||||||
|
nil,
|
||||||
|
newIp6AndPortFromString("[1::1]:6"),
|
||||||
|
nil,
|
||||||
|
}, alwaysAllowV6)
|
||||||
|
|
||||||
|
var cm *CacheMap
|
||||||
|
require.NotPanics(t, func() { cm = rl.CopyCache() })
|
||||||
|
|
||||||
|
c := (*cm)[owner.String()]
|
||||||
|
require.NotNil(t, c)
|
||||||
|
assert.ElementsMatch(t, []netip.AddrPort{
|
||||||
|
netip.MustParseAddrPort("1.2.3.4:5"),
|
||||||
|
netip.MustParseAddrPort("[1::1]:6"),
|
||||||
|
}, c.Reported)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRemoteList_Rebuild_SkipsNilReported drives unlockedCollect (via Rebuild) with
|
||||||
|
// nil reported entries and confirms only the valid addresses survive, with no panic.
|
||||||
|
func TestRemoteList_Rebuild_SkipsNilReported(t *testing.T) {
|
||||||
|
owner := netip.MustParseAddr("10.0.0.1")
|
||||||
|
rl := NewRemoteList([]netip.Addr{owner}, nil)
|
||||||
|
|
||||||
|
rl.unlockedSetV4(owner, owner, []*V4AddrPort{
|
||||||
|
nil,
|
||||||
|
newIp4AndPortFromString("1.2.3.4:5"),
|
||||||
|
}, alwaysAllowV4)
|
||||||
|
rl.unlockedSetV6(owner, owner, []*V6AddrPort{
|
||||||
|
newIp6AndPortFromString("[1::1]:6"),
|
||||||
|
nil,
|
||||||
|
}, alwaysAllowV6)
|
||||||
|
|
||||||
|
require.NotPanics(t, func() { rl.Rebuild([]netip.Prefix{}) })
|
||||||
|
|
||||||
|
assert.ElementsMatch(t, []netip.AddrPort{
|
||||||
|
netip.MustParseAddrPort("1.2.3.4:5"),
|
||||||
|
netip.MustParseAddrPort("[1::1]:6"),
|
||||||
|
}, rl.addrs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// newRelayControl marshals a NebulaControl the way it arrives on the wire so we can feed
|
||||||
|
// it through HandleControlMsg's unmarshal + validate path.
|
||||||
|
func newRelayControl(t *testing.T, typ NebulaControl_MessageType, from, to *Addr) []byte {
|
||||||
|
t.Helper()
|
||||||
|
msg := &NebulaControl{
|
||||||
|
Type: typ,
|
||||||
|
RelayFromAddr: from,
|
||||||
|
RelayToAddr: to,
|
||||||
|
}
|
||||||
|
b, err := msg.Marshal()
|
||||||
|
require.NoError(t, err)
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRelayManager_HandleControlMsg_NilRelayAddrs verifies the validation block added to
|
||||||
|
// HandleControlMsg: CreateRelay{Request,Response} carrying a nil RelayFromAddr or
|
||||||
|
// RelayToAddr are dropped with a debug log rather than nil-dereferencing downstream.
|
||||||
|
func TestRelayManager_HandleControlMsg_NilRelayAddrs(t *testing.T) {
|
||||||
|
good := netAddrToProtoAddr(netip.MustParseAddr("10.0.0.9"))
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
typ NebulaControl_MessageType
|
||||||
|
from *Addr
|
||||||
|
to *Addr
|
||||||
|
wantLog string // debug substring expected, "" == expect no drop log
|
||||||
|
}{
|
||||||
|
{"request nil from", NebulaControl_CreateRelayRequest, nil, good, "nil RelayFromAddr"},
|
||||||
|
{"request nil to", NebulaControl_CreateRelayRequest, good, nil, "nil RelayToAddr"},
|
||||||
|
{"request both nil", NebulaControl_CreateRelayRequest, nil, nil, "nil RelayFromAddr"},
|
||||||
|
{"response nil from", NebulaControl_CreateRelayResponse, nil, good, "nil RelayFromAddr"},
|
||||||
|
{"response nil to", NebulaControl_CreateRelayResponse, good, nil, "nil RelayToAddr"},
|
||||||
|
// A non-relay control type is not subject to the relay-addr validation and must
|
||||||
|
// pass through it untouched (the final switch simply no-ops on it).
|
||||||
|
{"unrelated type nil addrs", NebulaControl_None, nil, nil, ""},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
var buf bytes.Buffer
|
||||||
|
l := test.NewLoggerWithOutputAndLevel(&buf, slog.LevelDebug)
|
||||||
|
rm := &relayManager{l: l, hostmap: newHostMap(l)}
|
||||||
|
rm.useRelays.Store(true)
|
||||||
|
|
||||||
|
f := &Interface{l: l}
|
||||||
|
h := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("10.0.0.2")}, localIndexId: 1}
|
||||||
|
|
||||||
|
d := newRelayControl(t, tc.typ, tc.from, tc.to)
|
||||||
|
|
||||||
|
require.NotPanics(t, func() { rm.HandleControlMsg(h, d, f) })
|
||||||
|
|
||||||
|
if tc.wantLog == "" {
|
||||||
|
assert.NotContains(t, buf.String(), "nil Relay")
|
||||||
|
} else {
|
||||||
|
assert.Contains(t, buf.String(), tc.wantLog)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+20
-60
@@ -5,8 +5,6 @@ package nebula
|
|||||||
import (
|
import (
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
"github.com/google/gopacket"
|
|
||||||
"github.com/google/gopacket/layers"
|
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
@@ -22,7 +20,9 @@ func (c *Control) WaitForType(msgType header.MessageType, subType header.Message
|
|||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
pipeTo.InjectUDPPacket(p)
|
pipeTo.InjectUDPPacket(p)
|
||||||
if h.Type == msgType && h.Subtype == subType {
|
match := h.Type == msgType && h.Subtype == subType
|
||||||
|
p.Release()
|
||||||
|
if match {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -38,7 +38,9 @@ func (c *Control) WaitForTypeByIndex(toIndex uint32, msgType header.MessageType,
|
|||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
pipeTo.InjectUDPPacket(p)
|
pipeTo.InjectUDPPacket(p)
|
||||||
if h.RemoteIndex == toIndex && h.Type == msgType && h.Subtype == subType {
|
match := h.RemoteIndex == toIndex && h.Type == msgType && h.Subtype == subType
|
||||||
|
p.Release()
|
||||||
|
if match {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -90,65 +92,15 @@ func (c *Control) GetTunTxChan() <-chan []byte {
|
|||||||
return c.f.inside.(*overlay.TestTun).TxPackets
|
return c.f.inside.(*overlay.TestTun).TxPackets
|
||||||
}
|
}
|
||||||
|
|
||||||
// InjectUDPPacket will inject a packet into the udp side of nebula
|
// InjectUDPPacket injects a packet into the udp side. We copy internally so the caller keeps ownership of p.
|
||||||
|
// The copy comes from the freelist so steady-state alloc is zero.
|
||||||
func (c *Control) InjectUDPPacket(p *udp.Packet) {
|
func (c *Control) InjectUDPPacket(p *udp.Packet) {
|
||||||
c.f.outside.(*udp.TesterConn).Send(p)
|
c.f.outside.(*udp.TesterConn).Send(p.Copy())
|
||||||
}
|
}
|
||||||
|
|
||||||
// InjectTunUDPPacket puts a udp packet on the tun interface. Using UDP here because it's a simpler protocol
|
// InjectTunPacket pushes an IP packet onto the tun interface.
|
||||||
func (c *Control) InjectTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr netip.Addr, fromPort uint16, data []byte) {
|
func (c *Control) InjectTunPacket(packet []byte) {
|
||||||
serialize := make([]gopacket.SerializableLayer, 0)
|
c.f.inside.(*overlay.TestTun).Send(packet)
|
||||||
var netLayer gopacket.NetworkLayer
|
|
||||||
if toAddr.Is6() {
|
|
||||||
if !fromAddr.Is6() {
|
|
||||||
panic("Cant send ipv6 to ipv4")
|
|
||||||
}
|
|
||||||
ip := &layers.IPv6{
|
|
||||||
Version: 6,
|
|
||||||
NextHeader: layers.IPProtocolUDP,
|
|
||||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
|
||||||
DstIP: toAddr.Unmap().AsSlice(),
|
|
||||||
}
|
|
||||||
serialize = append(serialize, ip)
|
|
||||||
netLayer = ip
|
|
||||||
} else {
|
|
||||||
if !fromAddr.Is4() {
|
|
||||||
panic("Cant send ipv4 to ipv6")
|
|
||||||
}
|
|
||||||
|
|
||||||
ip := &layers.IPv4{
|
|
||||||
Version: 4,
|
|
||||||
TTL: 64,
|
|
||||||
Protocol: layers.IPProtocolUDP,
|
|
||||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
|
||||||
DstIP: toAddr.Unmap().AsSlice(),
|
|
||||||
}
|
|
||||||
serialize = append(serialize, ip)
|
|
||||||
netLayer = ip
|
|
||||||
}
|
|
||||||
|
|
||||||
udp := layers.UDP{
|
|
||||||
SrcPort: layers.UDPPort(fromPort),
|
|
||||||
DstPort: layers.UDPPort(toPort),
|
|
||||||
}
|
|
||||||
err := udp.SetNetworkLayerForChecksum(netLayer)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
buffer := gopacket.NewSerializeBuffer()
|
|
||||||
opt := gopacket.SerializeOptions{
|
|
||||||
ComputeChecksums: true,
|
|
||||||
FixLengths: true,
|
|
||||||
}
|
|
||||||
|
|
||||||
serialize = append(serialize, &udp, gopacket.Payload(data))
|
|
||||||
err = gopacket.SerializeLayers(buffer, opt, serialize...)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
c.f.inside.(*overlay.TestTun).Send(buffer.Bytes())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) GetVpnAddrs() []netip.Addr {
|
func (c *Control) GetVpnAddrs() []netip.Addr {
|
||||||
@@ -173,6 +125,14 @@ func (c *Control) GetHostmap() *HostMap {
|
|||||||
return c.f.hostMap
|
return c.f.hostMap
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetHostmapIndexCount returns the number of entries in the main hostmap Indexes table, holding
|
||||||
|
// the hostmap read lock so tests can poll it while connection manager churns tunnels.
|
||||||
|
func (c *Control) GetHostmapIndexCount() int {
|
||||||
|
c.f.hostMap.RLock()
|
||||||
|
defer c.f.hostMap.RUnlock()
|
||||||
|
return len(c.f.hostMap.Indexes)
|
||||||
|
}
|
||||||
|
|
||||||
func (c *Control) GetF() *Interface {
|
func (c *Control) GetF() *Interface {
|
||||||
return c.f
|
return c.f
|
||||||
}
|
}
|
||||||
|
|||||||
Vendored
+1
-1
@@ -62,7 +62,7 @@ function nebula.dissector(tvbuf, pktinfo, root)
|
|||||||
tree:add(pf_version, tvbuf:range(0,1))
|
tree:add(pf_version, tvbuf:range(0,1))
|
||||||
local type = tree:add(pf_type, tvbuf:range(0,1))
|
local type = tree:add(pf_type, tvbuf:range(0,1))
|
||||||
|
|
||||||
local nebula_type = bit32.band(tvbuf:range(0,1):uint(), 0x0F)
|
local nebula_type = bit.band(tvbuf:range(0,1):uint(), 0x0F)
|
||||||
if nebula_type == 0 then
|
if nebula_type == 0 then
|
||||||
local stage = tvbuf(8,8):uint64()
|
local stage = tvbuf(8,8):uint64()
|
||||||
tree:add(pf_subtype_handshake, tvbuf:range(1,1))
|
tree:add(pf_subtype_handshake, tvbuf:range(1,1))
|
||||||
|
|||||||
+92
-26
@@ -11,19 +11,21 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
type dnsServer struct {
|
type dnsServer struct {
|
||||||
sync.RWMutex
|
sync.RWMutex
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
dnsMap4 map[string]netip.Addr
|
dnsMap4 map[string]netip.Addr
|
||||||
dnsMap6 map[string]netip.Addr
|
dnsMap6 map[string]netip.Addr
|
||||||
hostMap *HostMap
|
hostMap *HostMap
|
||||||
myVpnAddrsTable *bart.Lite
|
pki *PKI
|
||||||
|
|
||||||
|
// selfHost is the cached FQDN we last seeded for ourselves
|
||||||
|
selfHost string
|
||||||
|
|
||||||
mux *dns.ServeMux
|
mux *dns.ServeMux
|
||||||
|
|
||||||
@@ -55,14 +57,14 @@ type dnsServer struct {
|
|||||||
// they no-op when DNS isn't enabled. Each Start invocation owns a ctx-cancel
|
// they no-op when DNS isn't enabled. Each Start invocation owns a ctx-cancel
|
||||||
// watcher that tears the listener down on nebula shutdown. The returned
|
// watcher that tears the listener down on nebula shutdown. The returned
|
||||||
// pointer is always non-nil, even on error.
|
// pointer is always non-nil, even on error.
|
||||||
func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, cs *CertState, hostMap *HostMap, c *config.C) (*dnsServer, error) {
|
func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, pki *PKI, hostMap *HostMap, c *config.C) (*dnsServer, error) {
|
||||||
ds := &dnsServer{
|
ds := &dnsServer{
|
||||||
l: l,
|
l: l,
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
dnsMap4: make(map[string]netip.Addr),
|
dnsMap4: make(map[string]netip.Addr),
|
||||||
dnsMap6: make(map[string]netip.Addr),
|
dnsMap6: make(map[string]netip.Addr),
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
myVpnAddrsTable: cs.myVpnAddrsTable,
|
pki: pki,
|
||||||
}
|
}
|
||||||
ds.mux = dns.NewServeMux()
|
ds.mux = dns.NewServeMux()
|
||||||
ds.mux.HandleFunc(".", ds.handleDnsRequest)
|
ds.mux.HandleFunc(".", ds.handleDnsRequest)
|
||||||
@@ -76,6 +78,7 @@ func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, cs *CertState,
|
|||||||
if err := ds.reload(c, true); err != nil {
|
if err := ds.reload(c, true); err != nil {
|
||||||
return ds, err
|
return ds, err
|
||||||
}
|
}
|
||||||
|
ds.seedSelf()
|
||||||
return ds, nil
|
return ds, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -113,7 +116,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
|||||||
d.Stop()
|
d.Stop()
|
||||||
}
|
}
|
||||||
// Drop any records that accumulated while enabled; a later re-enable
|
// Drop any records that accumulated while enabled; a later re-enable
|
||||||
// will repopulate from fresh handshakes.
|
// will repopulate from fresh handshakes and a fresh seedSelf.
|
||||||
d.clearRecords()
|
d.clearRecords()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -121,17 +124,14 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
|||||||
if running == nil {
|
if running == nil {
|
||||||
// Was disabled (or never started); bring it up now.
|
// Was disabled (or never started); bring it up now.
|
||||||
go d.Start()
|
go d.Start()
|
||||||
return nil
|
} else if !sameAddr {
|
||||||
|
d.shutdownServer(running, runningStarted, "reload")
|
||||||
|
// Old Start goroutine has now exited; bring up a fresh listener on the new address.
|
||||||
|
go d.Start()
|
||||||
}
|
}
|
||||||
|
|
||||||
if sameAddr {
|
// Refresh the self entry every enabled reload so cert renewals that change our name or VPN addresses are picked up.
|
||||||
return nil
|
d.seedSelf()
|
||||||
}
|
|
||||||
|
|
||||||
d.shutdownServer(running, runningStarted, "reload")
|
|
||||||
// Old Start goroutine has now exited; bring up a fresh listener on the
|
|
||||||
// new address.
|
|
||||||
go d.Start()
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -249,6 +249,20 @@ func (d *dnsServer) QueryCert(data string) string {
|
|||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// The hostmap only ever contains peers we have handshaked with, so it never carries an entry for ourselves.
|
||||||
|
// Answer self lookups straight from the local cert state.
|
||||||
|
if cs := d.certState(); cs != nil && cs.myVpnAddrsTable != nil && cs.myVpnAddrsTable.Contains(ip) {
|
||||||
|
c := cs.GetDefaultCertificate()
|
||||||
|
if c == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
b, err := c.MarshalJSON()
|
||||||
|
if err != nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return string(b)
|
||||||
|
}
|
||||||
|
|
||||||
hostinfo := d.hostMap.QueryVpnAddr(ip)
|
hostinfo := d.hostMap.QueryVpnAddr(ip)
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
return ""
|
return ""
|
||||||
@@ -266,12 +280,60 @@ func (d *dnsServer) QueryCert(data string) string {
|
|||||||
return string(b)
|
return string(b)
|
||||||
}
|
}
|
||||||
|
|
||||||
// clearRecords drops all DNS records.
|
// clearRecords drops all DNS records, including the self entry.
|
||||||
func (d *dnsServer) clearRecords() {
|
func (d *dnsServer) clearRecords() {
|
||||||
d.Lock()
|
d.Lock()
|
||||||
defer d.Unlock()
|
defer d.Unlock()
|
||||||
clear(d.dnsMap4)
|
clear(d.dnsMap4)
|
||||||
clear(d.dnsMap6)
|
clear(d.dnsMap6)
|
||||||
|
d.selfHost = ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// seedSelf inserts (or refreshes) a record for our own cert name pointing at our VPN addresses,
|
||||||
|
// so a single-lighthouse network can resolve the lighthouse's own hostname without the two-process workaround.
|
||||||
|
func (d *dnsServer) seedSelf() {
|
||||||
|
if !d.enabled.Load() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cs := d.certState()
|
||||||
|
if cs == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c := cs.GetDefaultCertificate()
|
||||||
|
if c == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
newHost := strings.ToLower(c.Name()) + "."
|
||||||
|
|
||||||
|
d.Lock()
|
||||||
|
defer d.Unlock()
|
||||||
|
if d.selfHost != "" && d.selfHost != newHost {
|
||||||
|
delete(d.dnsMap4, d.selfHost)
|
||||||
|
delete(d.dnsMap6, d.selfHost)
|
||||||
|
}
|
||||||
|
d.selfHost = newHost
|
||||||
|
delete(d.dnsMap4, newHost)
|
||||||
|
delete(d.dnsMap6, newHost)
|
||||||
|
haveV4, haveV6 := false, false
|
||||||
|
for _, addr := range cs.myVpnAddrs {
|
||||||
|
if addr.Is4() && !haveV4 {
|
||||||
|
d.dnsMap4[newHost] = addr
|
||||||
|
haveV4 = true
|
||||||
|
} else if addr.Is6() && !haveV6 {
|
||||||
|
d.dnsMap6[newHost] = addr
|
||||||
|
haveV6 = true
|
||||||
|
}
|
||||||
|
if haveV4 && haveV6 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *dnsServer) certState() *CertState {
|
||||||
|
if d.pki == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return d.pki.getCertState()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add adds the first IPv4 and IPv6 address that appears in `addresses` as the record for `host`
|
// Add adds the first IPv4 and IPv6 address that appears in `addresses` as the record for `host`
|
||||||
@@ -309,8 +371,12 @@ func (d *dnsServer) isSelfNebulaOrLocalhost(addr string) bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
cs := d.certState()
|
||||||
|
if cs == nil || cs.myVpnAddrsTable == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
//if we found it in this table, it's good
|
//if we found it in this table, it's good
|
||||||
return d.myVpnAddrsTable.Contains(b)
|
return cs.myVpnAddrsTable.Contains(b)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
||||||
|
|||||||
@@ -9,7 +9,10 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/gaissmai/bart"
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/cert_test"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -276,6 +279,92 @@ func TestDnsServer_Stop_beforeBind_doesNotHang(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// newTestPKI builds a minimal *PKI with a single v1 cert whose name and
|
||||||
|
// VPN addresses are caller-provided, suitable for exercising seedSelf and
|
||||||
|
// QueryCert self handling.
|
||||||
|
func newTestPKI(t *testing.T, name string, addrs []netip.Addr) *PKI {
|
||||||
|
t.Helper()
|
||||||
|
networks := make([]netip.Prefix, 0, len(addrs))
|
||||||
|
for _, a := range addrs {
|
||||||
|
bits := 32
|
||||||
|
if a.Is6() {
|
||||||
|
bits = 128
|
||||||
|
}
|
||||||
|
networks = append(networks, netip.PrefixFrom(a, bits))
|
||||||
|
}
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil)
|
||||||
|
c, _, _, _ := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, name, time.Time{}, time.Time{}, networks, nil, nil)
|
||||||
|
|
||||||
|
addrsTable := new(bart.Lite)
|
||||||
|
for _, a := range addrs {
|
||||||
|
addrsTable.Insert(netip.PrefixFrom(a, a.BitLen()))
|
||||||
|
}
|
||||||
|
|
||||||
|
cs := &CertState{
|
||||||
|
v2Cert: c,
|
||||||
|
initiatingVersion: cert.Version2,
|
||||||
|
myVpnAddrs: addrs,
|
||||||
|
myVpnAddrsTable: addrsTable,
|
||||||
|
}
|
||||||
|
pki := &PKI{}
|
||||||
|
pki.cs.Store(cs)
|
||||||
|
return pki
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDnsServer_seedSelf_addsOwnRecord(t *testing.T) {
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
myV4 := netip.MustParseAddr("10.0.0.1")
|
||||||
|
myV6 := netip.MustParseAddr("fd00::1")
|
||||||
|
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{myV4, myV6})
|
||||||
|
setDnsConfig(c, "127.0.0.1", "0", true, true)
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
|
||||||
|
ds.seedSelf()
|
||||||
|
got4, exists := ds.Query(dns.TypeA, "lighthouse.")
|
||||||
|
assert.True(t, exists)
|
||||||
|
assert.Equal(t, myV4, got4)
|
||||||
|
got6, exists := ds.Query(dns.TypeAAAA, "lighthouse.")
|
||||||
|
assert.True(t, exists)
|
||||||
|
assert.Equal(t, myV6, got6)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDnsServer_seedSelf_disabled_noOp(t *testing.T) {
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{netip.MustParseAddr("10.0.0.1")})
|
||||||
|
setDnsConfig(c, "127.0.0.1", "0", true, false)
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
|
||||||
|
ds.seedSelf()
|
||||||
|
_, exists := ds.Query(dns.TypeA, "lighthouse.")
|
||||||
|
assert.False(t, exists)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDnsServer_clearRecords_dropsSelfHost(t *testing.T) {
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{netip.MustParseAddr("10.0.0.1")})
|
||||||
|
setDnsConfig(c, "127.0.0.1", "0", true, true)
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
ds.seedSelf()
|
||||||
|
require.NotEmpty(t, ds.selfHost)
|
||||||
|
|
||||||
|
ds.clearRecords()
|
||||||
|
assert.Empty(t, ds.selfHost)
|
||||||
|
_, exists := ds.Query(dns.TypeA, "lighthouse.")
|
||||||
|
assert.False(t, exists)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDnsServer_QueryCert_returnsOwnCert(t *testing.T) {
|
||||||
|
ds, _ := newTestDnsServer(t)
|
||||||
|
myV4 := netip.MustParseAddr("10.0.0.1")
|
||||||
|
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{myV4})
|
||||||
|
|
||||||
|
got := ds.QueryCert(myV4.String() + ".")
|
||||||
|
assert.NotEmpty(t, got, "TXT lookup of our own VPN address should return our cert")
|
||||||
|
|
||||||
|
other := netip.MustParseAddr("10.0.0.99")
|
||||||
|
assert.Empty(t, ds.QueryCert(other.String()+"."), "unknown peer IP should return nothing")
|
||||||
|
}
|
||||||
|
|
||||||
func TestDnsServer_reload_disable_stopsRunningServer(t *testing.T) {
|
func TestDnsServer_reload_disable_stopsRunningServer(t *testing.T) {
|
||||||
port := freeUDPPort(t)
|
port := freeUDPPort(t)
|
||||||
ds, c := newTestDnsServer(t)
|
ds, c := newTestDnsServer(t)
|
||||||
|
|||||||
@@ -1,6 +1,16 @@
|
|||||||
FROM gcr.io/distroless/static:latest
|
FROM gcr.io/distroless/static:latest
|
||||||
|
|
||||||
ARG TARGETOS TARGETARCH
|
ARG TARGETOS TARGETARCH
|
||||||
|
|
||||||
|
ARG VERSION=dev
|
||||||
|
ARG REVISION=unknown
|
||||||
|
LABEL org.opencontainers.image.title="nebula" \
|
||||||
|
org.opencontainers.image.description="A scalable overlay networking tool with a focus on performance, simplicity and security" \
|
||||||
|
org.opencontainers.image.vendor="Nebula OSS" \
|
||||||
|
org.opencontainers.image.source="https://github.com/slackhq/nebula" \
|
||||||
|
org.opencontainers.image.version="${VERSION}" \
|
||||||
|
org.opencontainers.image.revision="${REVISION}"
|
||||||
|
|
||||||
COPY build/$TARGETOS-$TARGETARCH/nebula /nebula
|
COPY build/$TARGETOS-$TARGETARCH/nebula /nebula
|
||||||
COPY build/$TARGETOS-$TARGETARCH/nebula-cert /nebula-cert
|
COPY build/$TARGETOS-$TARGETARCH/nebula-cert /nebula-cert
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,85 @@
|
|||||||
|
//go:build e2e_testing
|
||||||
|
// +build e2e_testing
|
||||||
|
|
||||||
|
package e2e
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/cert_test"
|
||||||
|
"github.com/slackhq/nebula/e2e/router"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func assertTestRequestEchoed(t *testing.T, cipher string) {
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
over := m{"cipher": cipher}
|
||||||
|
a, aNet, aUdp, _ := newSimpleServer(cert.Version1, ca, caKey, "a", "10.128.0.1/24", over)
|
||||||
|
b, bNet, bUdp, _ := newSimpleServer(cert.Version1, ca, caKey, "b", "10.128.0.2/24", over)
|
||||||
|
|
||||||
|
a.InjectLightHouseAddr(bNet[0].Addr(), bUdp)
|
||||||
|
b.InjectLightHouseAddr(aNet[0].Addr(), aUdp)
|
||||||
|
a.Start()
|
||||||
|
b.Start()
|
||||||
|
t.Cleanup(func() { a.Stop(); b.Stop() })
|
||||||
|
r := router.NewR(t, a, b)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
|
||||||
|
assertTunnel(t, aNet[0].Addr(), bNet[0].Addr(), a, b, r)
|
||||||
|
drainUDPTx(a)
|
||||||
|
drainUDPTx(b)
|
||||||
|
|
||||||
|
payload := []byte("a test payload well over sixteen bytes long, wow it's so very long long long!")
|
||||||
|
require.Greater(t, len(payload), header.Len)
|
||||||
|
a.GetF().SendMessageToVpnAddr(header.Test, header.TestRequest, bNet[0].Addr(), payload, make([]byte, 12, 12), make([]byte, udp.MTU))
|
||||||
|
|
||||||
|
// Deliver A's request to B; B must echo a reply back
|
||||||
|
b.InjectUDPPacket(a.GetFromUDP(true))
|
||||||
|
reply := nextUDPTxOfType(t, b, header.Test, header.TestReply, 2*time.Second)
|
||||||
|
|
||||||
|
assert.Equal(t, aUdp, reply.To, "the reply must go back to the requester")
|
||||||
|
// header + echoed payload + 16-byte AEAD tag: proves the whole payload
|
||||||
|
// round-tripped rather than being dropped or truncated.
|
||||||
|
assert.Equal(t, header.Len+len(payload)+16, len(reply.Data), "the full payload must be echoed back")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTestRequestEchoesLongPayloadAES(t *testing.T) {
|
||||||
|
assertTestRequestEchoed(t, "aes")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTestRequestEchoesLongPayloadChaChaPoly(t *testing.T) {
|
||||||
|
assertTestRequestEchoed(t, "chachapoly")
|
||||||
|
}
|
||||||
|
|
||||||
|
// drainUDPTx empties a control's UDP tx queue without blocking.
|
||||||
|
func drainUDPTx(c *nebula.Control) {
|
||||||
|
for c.GetFromUDP(false) != nil {
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// nextUDPTxOfType returns the next packet a control transmits whose nebula
|
||||||
|
// header matches (wantType, wantSub), skipping unrelated packets.
|
||||||
|
// It fails the test if none arrives within the timeout.
|
||||||
|
func nextUDPTxOfType(t *testing.T, c *nebula.Control, wantType header.MessageType, wantSub header.MessageSubType, within time.Duration) *udp.Packet {
|
||||||
|
t.Helper()
|
||||||
|
ch := c.GetUDPTxChan()
|
||||||
|
timeout := time.After(within)
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case p := <-ch:
|
||||||
|
var h header.H
|
||||||
|
if err := h.Parse(p.Data); err == nil && h.Type == wantType && h.Subtype == wantSub {
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
case <-timeout:
|
||||||
|
t.Fatalf("timed out waiting for a %v/%v packet on the udp tx queue", wantType, wantSub)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -47,7 +47,7 @@ func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me to them")
|
t.Log("Trigger handshake from me to them")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
|
|
||||||
t.Log("Grab my msg1")
|
t.Log("Grab my msg1")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -97,7 +97,7 @@ func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake")
|
t.Log("Trigger handshake")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
|
|
||||||
t.Log("Get msg1 and deliver to responder")
|
t.Log("Get msg1 and deliver to responder")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -146,7 +146,7 @@ func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Complete a normal handshake")
|
t.Log("Complete a normal handshake")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
@@ -248,7 +248,7 @@ func TestHandshakeLateResponse(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me")
|
t.Log("Trigger handshake from me")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
|
|
||||||
t.Log("Grab msg1 but don't deliver")
|
t.Log("Grab msg1 but don't deliver")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -292,7 +292,7 @@ func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
|||||||
myControl.Start()
|
myControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me")
|
t.Log("Trigger handshake from me")
|
||||||
myControl.InjectTunUDPPacket(netip.MustParseAddr("10.128.0.2"), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(netip.MustParseAddr("10.128.0.2"), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
|
|
||||||
t.Log("Drain any handshake retransmits before injecting")
|
t.Log("Drain any handshake retransmits before injecting")
|
||||||
@@ -375,7 +375,7 @@ func TestHandshakeRemoteAllowList(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake from them")
|
t.Log("Trigger handshake from them")
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
msg1 := theirControl.GetFromUDP(true)
|
msg1 := theirControl.GetFromUDP(true)
|
||||||
|
|
||||||
t.Log("Rewrite the source to a blocked IP and inject")
|
t.Log("Rewrite the source to a blocked IP and inject")
|
||||||
@@ -426,7 +426,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Complete a normal handshake via the router")
|
t.Log("Complete a normal handshake via the router")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
@@ -437,7 +437,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|||||||
originalRemote := hi.CurrentRemote
|
originalRemote := hi.CurrentRemote
|
||||||
|
|
||||||
t.Log("Re-trigger traffic to cause a new handshake attempt (ErrAlreadySeen)")
|
t.Log("Re-trigger traffic to cause a new handshake attempt (ErrAlreadySeen)")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("roam"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("roam")))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
|
|
||||||
t.Log("Verify tunnel still works")
|
t.Log("Verify tunnel still works")
|
||||||
@@ -475,8 +475,8 @@ func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Send multiple packets to them (cached during handshake)")
|
t.Log("Send multiple packets to them (cached during handshake)")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet1"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet1")))
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet2"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet2")))
|
||||||
|
|
||||||
t.Log("Route until evil tunnel is closed")
|
t.Log("Route until evil tunnel is closed")
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
@@ -540,7 +540,7 @@ func TestHandshakeRelayComplete(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake via relay")
|
t.Log("Trigger handshake via relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi via relay"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi via relay")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi via relay"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi via relay"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
@@ -568,7 +568,7 @@ func TestHandshakeRelayComplete(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// NOTE: Relay V1 cert + IPv6 rejection is not tested here because
|
// NOTE: Relay V1 cert + IPv6 rejection is not tested here because
|
||||||
// InjectTunUDPPacket from a V4 node to a V6 address panics in the test
|
// BuildTunUDPPacket from a V4 node to a V6 address panics in the test
|
||||||
// framework. The check is in handshake_manager.go handleOutbound relay
|
// framework. The check is in handshake_manager.go handleOutbound relay
|
||||||
// logic (lines ~304-313): if the relay host has a V1 cert and either
|
// logic (lines ~304-313): if the relay host has a V1 cert and either
|
||||||
// address is IPv6, the relay is skipped.
|
// address is IPv6, the relay is skipped.
|
||||||
|
|||||||
+155
-60
@@ -16,6 +16,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/cert_test"
|
"github.com/slackhq/nebula/cert_test"
|
||||||
"github.com/slackhq/nebula/e2e/router"
|
"github.com/slackhq/nebula/e2e/router"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/overlay"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -39,11 +40,22 @@ func BenchmarkHotPath(b *testing.B) {
|
|||||||
r.CancelFlowLogs()
|
r.CancelFlowLogs()
|
||||||
|
|
||||||
assertTunnel(b, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(b, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
|
// Pre-build the IP packet bytes once so the bench measures the data plane,
|
||||||
|
// not gopacket SerializeLayers overhead.
|
||||||
|
prebuilt := BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
|
// EnableFanIn switches the router to a 0-alloc routing path. Required
|
||||||
|
// for hot-path benchmarks; would conflict with GetFromUDP-using tests.
|
||||||
|
r.EnableFanIn()
|
||||||
|
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
|
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(prebuilt)
|
||||||
_ = r.RouteForAllUntilTxTun(theirControl)
|
// Release the TUN-side bytes back to the harness freelist; the bench
|
||||||
|
// just confirms a packet arrived, the contents aren't inspected.
|
||||||
|
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
||||||
}
|
}
|
||||||
|
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
@@ -71,11 +83,15 @@ func BenchmarkHotPathRelay(b *testing.B) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
assertTunnel(b, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(b, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
|
|
||||||
|
prebuilt := BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
r.EnableFanIn()
|
||||||
|
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
|
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(prebuilt)
|
||||||
_ = r.RouteForAllUntilTxTun(theirControl)
|
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
||||||
}
|
}
|
||||||
|
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
@@ -97,7 +113,7 @@ func TestGoodHandshake(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||||
@@ -191,7 +207,7 @@ func TestWrongResponderHandshake(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
@@ -273,7 +289,7 @@ func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
@@ -352,8 +368,8 @@ func TestStage1Race(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake to start on both me and them")
|
t.Log("Trigger a handshake to start on both me and them")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
t.Log("Get both stage 1 handshake packets")
|
t.Log("Get both stage 1 handshake packets")
|
||||||
myHsForThem := myControl.GetFromUDP(true)
|
myHsForThem := myControl.GetFromUDP(true)
|
||||||
@@ -389,7 +405,7 @@ func TestStage1Race(t *testing.T) {
|
|||||||
|
|
||||||
r.Log("Spin until connection manager tears down a tunnel")
|
r.Log("Spin until connection manager tears down a tunnel")
|
||||||
|
|
||||||
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
|
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
t.Log("Connection manager hasn't ticked yet")
|
t.Log("Connection manager hasn't ticked yet")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
@@ -430,18 +446,20 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
r.Log("Trigger a handshake from me to them")
|
r.Log("Trigger a handshake from me to them")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
|
||||||
r.Log("Nuke my hostmap")
|
r.Log("Nuke my hostmap")
|
||||||
myHostmap := myControl.GetHostmap()
|
myHostmap := myControl.GetHostmap()
|
||||||
|
myHostmap.Lock()
|
||||||
myHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
myHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
||||||
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||||
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||||
|
myHostmap.Unlock()
|
||||||
|
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again")))
|
||||||
p = r.RouteForAllUntilTxTun(theirControl)
|
p = r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me again"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me again"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
|
||||||
@@ -449,10 +467,10 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
r.Log("Wait for the dead index to go away")
|
r.Log("Wait for the dead index to go away")
|
||||||
start := len(theirControl.GetHostmap().Indexes)
|
start := theirControl.GetHostmapIndexCount()
|
||||||
for {
|
for {
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
if len(theirControl.GetHostmap().Indexes) < start {
|
if theirControl.GetHostmapIndexCount() < start {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
@@ -480,7 +498,7 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
r.Log("Trigger a handshake from me to them")
|
r.Log("Trigger a handshake from me to them")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
@@ -488,11 +506,13 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
|
|
||||||
r.Log("Nuke my hostmap")
|
r.Log("Nuke my hostmap")
|
||||||
theirHostmap := theirControl.GetHostmap()
|
theirHostmap := theirControl.GetHostmap()
|
||||||
|
theirHostmap.Lock()
|
||||||
theirHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
theirHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
||||||
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||||
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||||
|
theirHostmap.Unlock()
|
||||||
|
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again")))
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
assertUdpPacket(t, []byte("Hi from them again"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from them again"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
||||||
r.RenderHostmaps("Derp hostmaps", myControl, theirControl)
|
r.RenderHostmaps("Derp hostmaps", myControl, theirControl)
|
||||||
@@ -501,10 +521,10 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
r.Log("Wait for the dead index to go away")
|
r.Log("Wait for the dead index to go away")
|
||||||
start := len(myControl.GetHostmap().Indexes)
|
start := myControl.GetHostmapIndexCount()
|
||||||
for {
|
for {
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
if len(myControl.GetHostmap().Indexes) < start {
|
if myControl.GetHostmapIndexCount() < start {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
@@ -535,7 +555,7 @@ func TestRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -565,7 +585,7 @@ func TestRelaysDontCareAboutIps(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -595,14 +615,14 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
|
||||||
t.Log("Ensure packet traversal from them to me via the relay")
|
t.Log("Ensure packet traversal from them to me via the relay")
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -612,12 +632,12 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
r.Log("Close the tunnel")
|
r.Log("Close the tunnel")
|
||||||
relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
|
relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
|
||||||
|
|
||||||
start := len(myControl.GetHostmap().Indexes)
|
start := myControl.GetHostmapIndexCount()
|
||||||
curIndexes := len(myControl.GetHostmap().Indexes)
|
curIndexes := myControl.GetHostmapIndexCount()
|
||||||
for curIndexes >= start {
|
for curIndexes >= start {
|
||||||
curIndexes = len(myControl.GetHostmap().Indexes)
|
curIndexes = myControl.GetHostmapIndexCount()
|
||||||
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail")))
|
||||||
|
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
return router.RouteAndExit
|
return router.RouteAndExit
|
||||||
@@ -634,7 +654,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(theirControl)
|
p = r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -669,7 +689,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
t.Log("Assert the tunnel works the other way, too")
|
t.Log("Assert the tunnel works the other way, too")
|
||||||
for {
|
for {
|
||||||
t.Log("RouteForAllUntilTxTun")
|
t.Log("RouteForAllUntilTxTun")
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -739,8 +759,8 @@ func TestStage1RaceRelays(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
||||||
|
|
||||||
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
r.Log("Wait for a packet from them to me")
|
r.Log("Wait for a packet from them to me")
|
||||||
p := r.RouteForAllUntilTxTun(myControl)
|
p := r.RouteForAllUntilTxTun(myControl)
|
||||||
@@ -787,8 +807,8 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
||||||
|
|
||||||
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
//r.RouteUntilAfterMsgType(myControl, header.Control, header.MessageNone)
|
//r.RouteUntilAfterMsgType(myControl, header.Control, header.MessageNone)
|
||||||
//r.RouteUntilAfterMsgType(theirControl, header.Control, header.MessageNone)
|
//r.RouteUntilAfterMsgType(theirControl, header.Control, header.MessageNone)
|
||||||
@@ -803,18 +823,18 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
|
|
||||||
t.Log("Wait until we remove extra tunnels")
|
t.Log("Wait until we remove extra tunnels")
|
||||||
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
||||||
len(myControl.GetHostmap().Indexes),
|
myControl.GetHostmapIndexCount(),
|
||||||
len(theirControl.GetHostmap().Indexes),
|
theirControl.GetHostmapIndexCount(),
|
||||||
len(relayControl.GetHostmap().Indexes),
|
relayControl.GetHostmapIndexCount(),
|
||||||
)
|
)
|
||||||
hostInfos := len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
|
hostInfos := myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount()
|
||||||
retries := 60
|
retries := 60
|
||||||
for hostInfos > 6 && retries > 0 {
|
for hostInfos > 6 && retries > 0 {
|
||||||
hostInfos = len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
|
hostInfos = myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount()
|
||||||
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
||||||
len(myControl.GetHostmap().Indexes),
|
myControl.GetHostmapIndexCount(),
|
||||||
len(theirControl.GetHostmap().Indexes),
|
theirControl.GetHostmapIndexCount(),
|
||||||
len(relayControl.GetHostmap().Indexes),
|
relayControl.GetHostmapIndexCount(),
|
||||||
)
|
)
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
t.Log("Connection manager hasn't ticked yet")
|
t.Log("Connection manager hasn't ticked yet")
|
||||||
@@ -852,7 +872,7 @@ func TestRehandshakingRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -908,24 +928,24 @@ func TestRehandshakingRelays(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
|
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
|
||||||
// We should have two hostinfos on all sides
|
// We should have two hostinfos on all sides
|
||||||
for len(myControl.GetHostmap().Indexes) != 2 {
|
for myControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
|
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
}
|
}
|
||||||
t.Logf("myControl hostinfos got cleaned up!")
|
t.Logf("myControl hostinfos got cleaned up!")
|
||||||
for len(theirControl.GetHostmap().Indexes) != 2 {
|
for theirControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
|
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
}
|
}
|
||||||
t.Logf("theirControl hostinfos got cleaned up!")
|
t.Logf("theirControl hostinfos got cleaned up!")
|
||||||
for len(relayControl.GetHostmap().Indexes) != 2 {
|
for relayControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
|
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
@@ -957,7 +977,7 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -1013,24 +1033,24 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
|
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
|
||||||
// We should have two hostinfos on all sides
|
// We should have two hostinfos on all sides
|
||||||
for len(myControl.GetHostmap().Indexes) != 2 {
|
for myControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
|
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
}
|
}
|
||||||
t.Logf("myControl hostinfos got cleaned up!")
|
t.Logf("myControl hostinfos got cleaned up!")
|
||||||
for len(theirControl.GetHostmap().Indexes) != 2 {
|
for theirControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
|
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
}
|
}
|
||||||
t.Logf("theirControl hostinfos got cleaned up!")
|
t.Logf("theirControl hostinfos got cleaned up!")
|
||||||
for len(relayControl.GetHostmap().Indexes) != 2 {
|
for relayControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
|
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
@@ -1107,7 +1127,7 @@ func TestRehandshaking(t *testing.T) {
|
|||||||
theirConfig.ReloadConfigString(string(rc))
|
theirConfig.ReloadConfigString(string(rc))
|
||||||
|
|
||||||
r.Log("Spin until there is only 1 tunnel")
|
r.Log("Spin until there is only 1 tunnel")
|
||||||
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
|
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
t.Log("Connection manager hasn't ticked yet")
|
t.Log("Connection manager hasn't ticked yet")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
@@ -1207,7 +1227,7 @@ func TestRehandshakingLoser(t *testing.T) {
|
|||||||
myConfig.ReloadConfigString(string(rc))
|
myConfig.ReloadConfigString(string(rc))
|
||||||
|
|
||||||
r.Log("Spin until there is only 1 tunnel")
|
r.Log("Spin until there is only 1 tunnel")
|
||||||
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
|
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
t.Log("Connection manager hasn't ticked yet")
|
t.Log("Connection manager hasn't ticked yet")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
@@ -1259,8 +1279,8 @@ func TestRaceRegression(t *testing.T) {
|
|||||||
//them rx stage:2 initiatorIndex=120607833 responderIndex=4209862089
|
//them rx stage:2 initiatorIndex=120607833 responderIndex=4209862089
|
||||||
|
|
||||||
t.Log("Start both handshakes")
|
t.Log("Start both handshakes")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
t.Log("Get both stage 1")
|
t.Log("Get both stage 1")
|
||||||
myStage1ForThem := myControl.GetFromUDP(true)
|
myStage1ForThem := myControl.GetFromUDP(true)
|
||||||
@@ -1476,7 +1496,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||||
myControl.InjectTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||||
@@ -1504,7 +1524,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
|||||||
assertUdpPacket(t, []byte("Hi from me"), myCachedPacket, myVpnIpNet[0].Addr(), spookyDest, 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), myCachedPacket, myVpnIpNet[0].Addr(), spookyDest, 80, 80)
|
||||||
|
|
||||||
//reply
|
//reply
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman")))
|
||||||
//wait for reply
|
//wait for reply
|
||||||
theirControl.WaitForType(1, 0, myControl)
|
theirControl.WaitForType(1, 0, myControl)
|
||||||
theirCachedPacket := myControl.GetFromTun(true)
|
theirCachedPacket := myControl.GetFromTun(true)
|
||||||
@@ -1519,3 +1539,78 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
|||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
theirControl.Stop()
|
theirControl.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMultiVpnAddrDeletePrimaryKeepsSecondAddr(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
// Regression for the hostmap multi-vpnAddr delete bug. A dual-stack (v4+v6) V2-cert peer that
|
||||||
|
// handshakes twice at once ends up with two hostinfos linked in the shared next/prev chain, with the
|
||||||
|
// primary owning both addresses. Deleting that primary (e.g. connection manager dropping it, a
|
||||||
|
// CloseTunnel, a collision) must promote the surviving sibling for EVERY address. The pre-fix code
|
||||||
|
// unlinked the chain once per address, so it promoted the sibling for the first address and orphaned
|
||||||
|
// the second: the peer stayed reachable at its v4 addr but not its v6 addr despite a live tunnel.
|
||||||
|
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24,fd00::1/64", nil)
|
||||||
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.2/24,fd00::2/64", nil)
|
||||||
|
|
||||||
|
// This bug only exists for peers carrying more than one vpn address
|
||||||
|
require.Len(t, theirVpnIpNet, 2)
|
||||||
|
theirV4 := theirVpnIpNet[0].Addr()
|
||||||
|
theirV6 := theirVpnIpNet[1].Addr()
|
||||||
|
|
||||||
|
// Put their info in our lighthouse and vice versa
|
||||||
|
myControl.InjectLightHouseAddr(theirV4, theirUdpAddr)
|
||||||
|
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||||
|
|
||||||
|
// Build a router so we don't have to reason who gets which packet
|
||||||
|
r := router.NewR(t, myControl, theirControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
|
||||||
|
myControl.Start()
|
||||||
|
theirControl.Start()
|
||||||
|
|
||||||
|
// Race a handshake so both of us build a hostinfo for the other, leaving my hostmap with a single
|
||||||
|
// host (them) backed by two linked hostinfos, just like TestStage1Race.
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirV4, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirV4, 80, []byte("Hi from them")))
|
||||||
|
|
||||||
|
myHsForThem := myControl.GetFromUDP(true)
|
||||||
|
theirHsForMe := theirControl.GetFromUDP(true)
|
||||||
|
|
||||||
|
r.InjectUDPPacket(theirControl, myControl, theirHsForMe)
|
||||||
|
r.InjectUDPPacket(myControl, theirControl, myHsForThem)
|
||||||
|
|
||||||
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
|
r.RouteForAllUntilTxTun(myControl)
|
||||||
|
|
||||||
|
r.RenderHostmaps("Racing hostmaps", myControl, theirControl)
|
||||||
|
|
||||||
|
// Two hostinfos for them means the shared next/prev chain has a sibling to promote. The Hosts map has
|
||||||
|
// one entry per vpn address (two, for dual stack), so the index count is what tells us there are two
|
||||||
|
// hostinfos.
|
||||||
|
require.Len(t, myControl.ListHostmapIndexes(false), 2)
|
||||||
|
|
||||||
|
// The primary owns both of their addresses
|
||||||
|
primaryV4 := myControl.GetHostInfoByVpnAddr(theirV4, false)
|
||||||
|
primaryV6 := myControl.GetHostInfoByVpnAddr(theirV6, false)
|
||||||
|
require.NotNil(t, primaryV4)
|
||||||
|
require.NotNil(t, primaryV6)
|
||||||
|
require.Equal(t, primaryV4.LocalIndex, primaryV6.LocalIndex, "both addrs should point at the same primary")
|
||||||
|
|
||||||
|
// Delete the primary tunnel. localOnly so we don't perturb their side, we only care about my hostmap.
|
||||||
|
require.True(t, myControl.CloseTunnel(theirV4, true))
|
||||||
|
|
||||||
|
// The surviving sibling must still serve BOTH addresses.
|
||||||
|
survivorV4 := myControl.GetHostInfoByVpnAddr(theirV4, false)
|
||||||
|
survivorV6 := myControl.GetHostInfoByVpnAddr(theirV6, false)
|
||||||
|
require.NotNil(t, survivorV4, "v4 addr should still resolve to the surviving tunnel")
|
||||||
|
// Pre-fix this is nil: the second address was orphaned when the primary was deleted.
|
||||||
|
require.NotNil(t, survivorV6, "v6 addr was orphaned after deleting the primary (multi-vpnAddr delete bug)")
|
||||||
|
assert.Equal(t, survivorV4.LocalIndex, survivorV6.LocalIndex, "both addrs should promote to the same survivor")
|
||||||
|
assert.NotEqual(t, primaryV4.LocalIndex, survivorV4.LocalIndex, "a different hostinfo should now be primary")
|
||||||
|
|
||||||
|
r.RenderHostmaps("Final hostmaps", myControl, theirControl)
|
||||||
|
|
||||||
|
myControl.Stop()
|
||||||
|
theirControl.Stop()
|
||||||
|
}
|
||||||
|
|||||||
+59
-6
@@ -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"
|
||||||
@@ -294,12 +292,12 @@ func deadline(t *testing.T, seconds time.Duration) doneCb {
|
|||||||
|
|
||||||
func assertTunnel(t testing.TB, vpnIpA, vpnIpB netip.Addr, controlA, controlB *nebula.Control, r *router.R) {
|
func assertTunnel(t testing.TB, vpnIpA, vpnIpB netip.Addr, controlA, controlB *nebula.Control, r *router.R) {
|
||||||
// Send a packet from them to me
|
// Send a packet from them to me
|
||||||
controlB.InjectTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B"))
|
controlB.InjectTunPacket(BuildTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B")))
|
||||||
bPacket := r.RouteForAllUntilTxTun(controlA)
|
bPacket := r.RouteForAllUntilTxTun(controlA)
|
||||||
assertUdpPacket(t, []byte("Hi from B"), bPacket, vpnIpB, vpnIpA, 90, 80)
|
assertUdpPacket(t, []byte("Hi from B"), bPacket, vpnIpB, vpnIpA, 90, 80)
|
||||||
|
|
||||||
// And once more from me to them
|
// And once more from me to them
|
||||||
controlA.InjectTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A"))
|
controlA.InjectTunPacket(BuildTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A")))
|
||||||
aPacket := r.RouteForAllUntilTxTun(controlB)
|
aPacket := r.RouteForAllUntilTxTun(controlB)
|
||||||
assertUdpPacket(t, []byte("Hello from A"), aPacket, vpnIpA, vpnIpB, 90, 80)
|
assertUdpPacket(t, []byte("Hello from A"), aPacket, vpnIpA, vpnIpB, 90, 80)
|
||||||
}
|
}
|
||||||
@@ -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
|
||||||
@@ -408,3 +406,58 @@ func testLogLevelName() string {
|
|||||||
}
|
}
|
||||||
return "info"
|
return "info"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// BuildTunUDPPacket assembles an IP+UDP packet suitable for Control.InjectTunPacket.
|
||||||
|
// Using UDP here because it's a simpler protocol.
|
||||||
|
func BuildTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr netip.Addr, fromPort uint16, data []byte) []byte {
|
||||||
|
serialize := make([]gopacket.SerializableLayer, 0)
|
||||||
|
var netLayer gopacket.NetworkLayer
|
||||||
|
if toAddr.Is6() {
|
||||||
|
if !fromAddr.Is6() {
|
||||||
|
panic("Cant send ipv6 to ipv4")
|
||||||
|
}
|
||||||
|
ip := &layers.IPv6{
|
||||||
|
Version: 6,
|
||||||
|
NextHeader: layers.IPProtocolUDP,
|
||||||
|
SrcIP: fromAddr.Unmap().AsSlice(),
|
||||||
|
DstIP: toAddr.Unmap().AsSlice(),
|
||||||
|
}
|
||||||
|
serialize = append(serialize, ip)
|
||||||
|
netLayer = ip
|
||||||
|
} else {
|
||||||
|
if !fromAddr.Is4() {
|
||||||
|
panic("Cant send ipv4 to ipv6")
|
||||||
|
}
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Protocol: layers.IPProtocolUDP,
|
||||||
|
SrcIP: fromAddr.Unmap().AsSlice(),
|
||||||
|
DstIP: toAddr.Unmap().AsSlice(),
|
||||||
|
}
|
||||||
|
serialize = append(serialize, ip)
|
||||||
|
netLayer = ip
|
||||||
|
}
|
||||||
|
|
||||||
|
udp := layers.UDP{
|
||||||
|
SrcPort: layers.UDPPort(fromPort),
|
||||||
|
DstPort: layers.UDPPort(toPort),
|
||||||
|
}
|
||||||
|
if err := udp.SetNetworkLayerForChecksum(netLayer); err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
buffer := gopacket.NewSerializeBuffer()
|
||||||
|
opt := gopacket.SerializeOptions{
|
||||||
|
ComputeChecksums: true,
|
||||||
|
FixLengths: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
serialize = append(serialize, &udp, gopacket.Payload(data))
|
||||||
|
if err := gopacket.SerializeLayers(buffer, opt, serialize...); err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return buffer.Bytes()
|
||||||
|
}
|
||||||
|
|||||||
+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)
|
||||||
|
|||||||
+188
-54
@@ -13,6 +13,7 @@ import (
|
|||||||
"regexp"
|
"regexp"
|
||||||
"sort"
|
"sort"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -24,6 +25,19 @@ import (
|
|||||||
"golang.org/x/exp/maps"
|
"golang.org/x/exp/maps"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// outNatKey is the (from, to) pair used by outNat. Comparable struct, so it works as a map key without the
|
||||||
|
// allocation cost of a string-concat key.
|
||||||
|
type outNatKey struct {
|
||||||
|
from, to netip.AddrPort
|
||||||
|
}
|
||||||
|
|
||||||
|
// fannedPacket pairs a UDP TX packet with its source control so the router can route it after popping from
|
||||||
|
// the fan-in channel.
|
||||||
|
type fannedPacket struct {
|
||||||
|
from *nebula.Control
|
||||||
|
pkt *udp.Packet
|
||||||
|
}
|
||||||
|
|
||||||
type R struct {
|
type R struct {
|
||||||
// Simple map of the ip:port registered on a control to the control
|
// Simple map of the ip:port registered on a control to the control
|
||||||
// Basically a router, right?
|
// Basically a router, right?
|
||||||
@@ -34,12 +48,28 @@ type R struct {
|
|||||||
|
|
||||||
// A last used map, if an inbound packet hit the inNat map then
|
// A last used map, if an inbound packet hit the inNat map then
|
||||||
// all return packets should use the same last used inbound address for the outbound sender
|
// all return packets should use the same last used inbound address for the outbound sender
|
||||||
// map[from address + ":" + to address] => ip:port to rewrite in the udp packet to receiver
|
outNat map[outNatKey]netip.AddrPort
|
||||||
outNat map[string]netip.AddrPort
|
|
||||||
|
|
||||||
// A map of vpn ip to the nebula control it belongs to
|
// A map of vpn ip to the nebula control it belongs to
|
||||||
vpnControls map[netip.Addr]*nebula.Control
|
vpnControls map[netip.Addr]*nebula.Control
|
||||||
|
|
||||||
|
// Cached select infrastructure for RouteForAllUntilTxTun.
|
||||||
|
// The controls map is immutable after NewR so the cases are good for the test lifetime.
|
||||||
|
// We only rebuild if a different receiver is asked.
|
||||||
|
selRecvCtl *nebula.Control
|
||||||
|
selCases []reflect.SelectCase
|
||||||
|
selCtls []*nebula.Control
|
||||||
|
|
||||||
|
// Optional fan-in mode for hot-path benchmarks: one forwarder goroutine per control drains UDP TX into udpFanIn,
|
||||||
|
// so RouteForAllUntilTxTun can do a fixed 2-way native select instead of paying reflect.Select per call.
|
||||||
|
// Off by default (would otherwise interleave with tests that use GetFromUDP directly on the same control).
|
||||||
|
// Enabled by EnableFanIn.
|
||||||
|
udpFanIn chan fannedPacket
|
||||||
|
stopFanIn chan struct{}
|
||||||
|
fanInWG sync.WaitGroup
|
||||||
|
fanInMu sync.Mutex
|
||||||
|
fanInOn atomic.Bool
|
||||||
|
|
||||||
ignoreFlows []ignoreFlow
|
ignoreFlows []ignoreFlow
|
||||||
flow []flowEntry
|
flow []flowEntry
|
||||||
|
|
||||||
@@ -119,7 +149,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
|||||||
controls: make(map[netip.AddrPort]*nebula.Control),
|
controls: make(map[netip.AddrPort]*nebula.Control),
|
||||||
vpnControls: make(map[netip.Addr]*nebula.Control),
|
vpnControls: make(map[netip.Addr]*nebula.Control),
|
||||||
inNat: make(map[netip.AddrPort]*nebula.Control),
|
inNat: make(map[netip.AddrPort]*nebula.Control),
|
||||||
outNat: make(map[string]netip.AddrPort),
|
outNat: make(map[outNatKey]netip.AddrPort),
|
||||||
flow: []flowEntry{},
|
flow: []flowEntry{},
|
||||||
ignoreFlows: []ignoreFlow{},
|
ignoreFlows: []ignoreFlow{},
|
||||||
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
||||||
@@ -153,8 +183,10 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
|||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return
|
||||||
case <-clockSource.C:
|
case <-clockSource.C:
|
||||||
|
r.Lock()
|
||||||
r.renderHostmaps("clock tick")
|
r.renderHostmaps("clock tick")
|
||||||
r.renderFlow()
|
r.renderFlow()
|
||||||
|
r.Unlock()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
@@ -180,15 +212,21 @@ func (r *R) AddRoute(ip netip.Addr, port uint16, c *nebula.Control) {
|
|||||||
// RenderFlow renders the packet flow seen up until now and stops further automatic renders from happening.
|
// RenderFlow renders the packet flow seen up until now and stops further automatic renders from happening.
|
||||||
func (r *R) RenderFlow() {
|
func (r *R) RenderFlow() {
|
||||||
r.cancelRender()
|
r.cancelRender()
|
||||||
|
r.Lock()
|
||||||
|
defer r.Unlock()
|
||||||
r.renderFlow()
|
r.renderFlow()
|
||||||
}
|
}
|
||||||
|
|
||||||
// CancelFlowLogs stops flow logs from being tracked and destroys any logs already collected
|
// CancelFlowLogs stops flow logs from being tracked and destroys any logs already collected
|
||||||
func (r *R) CancelFlowLogs() {
|
func (r *R) CancelFlowLogs() {
|
||||||
r.cancelRender()
|
r.cancelRender()
|
||||||
|
r.Lock()
|
||||||
r.flow = nil
|
r.flow = nil
|
||||||
|
r.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// renderFlow writes the flow log to disk. Caller must hold r.Lock. renderFlow reads r.flow / r.additionalGraphs and
|
||||||
|
// the *packet pointers stashed inside, all of which are mutated under the same lock by routing paths.
|
||||||
func (r *R) renderFlow() {
|
func (r *R) renderFlow() {
|
||||||
if r.flow == nil {
|
if r.flow == nil {
|
||||||
return
|
return
|
||||||
@@ -434,68 +472,157 @@ func (r *R) RouteUntilTxTun(sender *nebula.Control, receiver *nebula.Control) []
|
|||||||
panic("No control for udp tx " + a.String())
|
panic("No control for udp tx " + a.String())
|
||||||
}
|
}
|
||||||
fp := r.unlockedInjectFlow(sender, c, p, false)
|
fp := r.unlockedInjectFlow(sender, c, p, false)
|
||||||
c.InjectUDPPacket(p)
|
c.InjectUDPPacket(p) // copies internally; original is ours to release
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// RouteForAllUntilTxTun will route for everyone and return when a packet is seen on receivers tun
|
// RouteForAllUntilTxTun will route for everyone and return when a packet is seen on the receiver's tun.
|
||||||
// If the router doesn't have the nebula controller for that address, we panic
|
// If a control's UDP TX address can't be matched to a registered control, we panic.
|
||||||
|
//
|
||||||
|
// For allocation-sensitive callers (hot-path benchmarks, in particular relay
|
||||||
|
// benches with 3+ controls), call EnableFanIn() first.
|
||||||
func (r *R) RouteForAllUntilTxTun(receiver *nebula.Control) []byte {
|
func (r *R) RouteForAllUntilTxTun(receiver *nebula.Control) []byte {
|
||||||
|
if r.fanInOn.Load() {
|
||||||
|
return r.routeFanIn(receiver)
|
||||||
|
}
|
||||||
|
return r.routeReflect(receiver)
|
||||||
|
}
|
||||||
|
|
||||||
|
// routeFanIn is the alloc-free path used when EnableFanIn is in effect.
|
||||||
|
func (r *R) routeFanIn(receiver *nebula.Control) []byte {
|
||||||
|
tunTx := receiver.GetTunTxChan()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case p := <-tunTx:
|
||||||
|
r.Lock()
|
||||||
|
if r.flow != nil {
|
||||||
|
np := udp.Packet{Data: make([]byte, len(p))}
|
||||||
|
copy(np.Data, p)
|
||||||
|
r.unlockedInjectFlow(receiver, receiver, &np, true)
|
||||||
|
}
|
||||||
|
r.Unlock()
|
||||||
|
return p
|
||||||
|
case fp := <-r.udpFanIn:
|
||||||
|
r.routeUDP(fp.from, fp.pkt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// routeReflect is the default reflect.Select-based path. Pays the boxing allocation per call but doesn't interfere
|
||||||
|
// with tests that pull packets directly from controls' UDP TX channels via GetFromUDP.
|
||||||
|
func (r *R) routeReflect(receiver *nebula.Control) []byte {
|
||||||
|
sc, cm := r.selectCasesFor(receiver)
|
||||||
|
for {
|
||||||
|
x, rx, _ := reflect.Select(sc)
|
||||||
|
if x == 0 {
|
||||||
|
p := rx.Interface().([]byte)
|
||||||
|
r.Lock()
|
||||||
|
if r.flow != nil {
|
||||||
|
np := udp.Packet{Data: make([]byte, len(p))}
|
||||||
|
copy(np.Data, p)
|
||||||
|
r.unlockedInjectFlow(cm[x], cm[x], &np, true)
|
||||||
|
}
|
||||||
|
r.Unlock()
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
r.routeUDP(cm[x], rx.Interface().(*udp.Packet))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// EnableFanIn switches RouteForAllUntilTxTun to the alloc-free fan-in path.
|
||||||
|
// One forwarder goroutine per registered control drains UDP TX into a shared channel that RouteForAllUntilTxTun selects
|
||||||
|
// on alongside the receiver's TUN TX channel.
|
||||||
|
func (r *R) EnableFanIn() {
|
||||||
|
r.fanInMu.Lock()
|
||||||
|
defer r.fanInMu.Unlock()
|
||||||
|
if r.fanInOn.Load() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.udpFanIn = make(chan fannedPacket, 32)
|
||||||
|
r.stopFanIn = make(chan struct{})
|
||||||
|
for _, c := range r.controls {
|
||||||
|
r.startFanInWorker(c)
|
||||||
|
}
|
||||||
|
r.fanInOn.Store(true)
|
||||||
|
r.t.Cleanup(r.stopFanInWorkers)
|
||||||
|
}
|
||||||
|
|
||||||
|
// startFanInWorker spawns a goroutine that drains c's UDP TX into r.udpFanIn.
|
||||||
|
func (r *R) startFanInWorker(c *nebula.Control) {
|
||||||
|
r.fanInWG.Add(1)
|
||||||
|
udpTx := c.GetUDPTxChan()
|
||||||
|
go func() {
|
||||||
|
defer r.fanInWG.Done()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-r.stopFanIn:
|
||||||
|
return
|
||||||
|
case p := <-udpTx:
|
||||||
|
select {
|
||||||
|
case <-r.stopFanIn:
|
||||||
|
p.Release()
|
||||||
|
return
|
||||||
|
case r.udpFanIn <- fannedPacket{from: c, pkt: p}:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// stopFanInWorkers signals the fan-in goroutines to exit and waits for them.
|
||||||
|
func (r *R) stopFanInWorkers() {
|
||||||
|
r.fanInMu.Lock()
|
||||||
|
wasOn := r.fanInOn.Swap(false)
|
||||||
|
r.fanInMu.Unlock()
|
||||||
|
if !wasOn {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
close(r.stopFanIn)
|
||||||
|
r.fanInWG.Wait()
|
||||||
|
}
|
||||||
|
|
||||||
|
// routeUDP forwards a UDP TX packet from the named source control to the destination control derived from p.To,
|
||||||
|
// releasing the source packet after InjectUDPPacket has copied its bytes into a fresh pool slot.
|
||||||
|
func (r *R) routeUDP(from *nebula.Control, p *udp.Packet) {
|
||||||
|
r.Lock()
|
||||||
|
defer r.Unlock()
|
||||||
|
a := from.GetUDPAddr()
|
||||||
|
c := r.getControl(a, p.To, p)
|
||||||
|
if c == nil {
|
||||||
|
panic(fmt.Sprintf("No control for udp tx %s", p.To))
|
||||||
|
}
|
||||||
|
fp := r.unlockedInjectFlow(from, c, p, false)
|
||||||
|
c.InjectUDPPacket(p) // copies internally; original is ours to release
|
||||||
|
fp.WasReceived()
|
||||||
|
p.Release()
|
||||||
|
}
|
||||||
|
|
||||||
|
// selectCasesFor returns the SelectCase array used by routeReflect: one slot for the receiver's TUN TX channel followed
|
||||||
|
// by one per control's UDP TX channel. Cached for the test lifetime, only rebuilt if the receiver changes.
|
||||||
|
func (r *R) selectCasesFor(receiver *nebula.Control) ([]reflect.SelectCase, []*nebula.Control) {
|
||||||
|
r.Lock()
|
||||||
|
defer r.Unlock()
|
||||||
|
if r.selRecvCtl == receiver && r.selCases != nil {
|
||||||
|
return r.selCases, r.selCtls
|
||||||
|
}
|
||||||
sc := make([]reflect.SelectCase, len(r.controls)+1)
|
sc := make([]reflect.SelectCase, len(r.controls)+1)
|
||||||
cm := make([]*nebula.Control, len(r.controls)+1)
|
cm := make([]*nebula.Control, len(r.controls)+1)
|
||||||
|
sc[0] = reflect.SelectCase{Dir: reflect.SelectRecv, Chan: reflect.ValueOf(receiver.GetTunTxChan())}
|
||||||
i := 0
|
cm[0] = receiver
|
||||||
sc[i] = reflect.SelectCase{
|
i := 1
|
||||||
Dir: reflect.SelectRecv,
|
|
||||||
Chan: reflect.ValueOf(receiver.GetTunTxChan()),
|
|
||||||
Send: reflect.Value{},
|
|
||||||
}
|
|
||||||
cm[i] = receiver
|
|
||||||
|
|
||||||
i++
|
|
||||||
for _, c := range r.controls {
|
for _, c := range r.controls {
|
||||||
sc[i] = reflect.SelectCase{
|
sc[i] = reflect.SelectCase{Dir: reflect.SelectRecv, Chan: reflect.ValueOf(c.GetUDPTxChan())}
|
||||||
Dir: reflect.SelectRecv,
|
|
||||||
Chan: reflect.ValueOf(c.GetUDPTxChan()),
|
|
||||||
Send: reflect.Value{},
|
|
||||||
}
|
|
||||||
|
|
||||||
cm[i] = c
|
cm[i] = c
|
||||||
i++
|
i++
|
||||||
}
|
}
|
||||||
|
r.selRecvCtl = receiver
|
||||||
for {
|
r.selCases = sc
|
||||||
x, rx, _ := reflect.Select(sc)
|
r.selCtls = cm
|
||||||
r.Lock()
|
return sc, cm
|
||||||
|
|
||||||
if x == 0 {
|
|
||||||
// we are the tun tx, we can exit
|
|
||||||
p := rx.Interface().([]byte)
|
|
||||||
np := udp.Packet{Data: make([]byte, len(p))}
|
|
||||||
copy(np.Data, p)
|
|
||||||
|
|
||||||
r.unlockedInjectFlow(cm[x], cm[x], &np, true)
|
|
||||||
r.Unlock()
|
|
||||||
return p
|
|
||||||
|
|
||||||
} else {
|
|
||||||
// we are a udp tx, route and continue
|
|
||||||
p := rx.Interface().(*udp.Packet)
|
|
||||||
a := cm[x].GetUDPAddr()
|
|
||||||
c := r.getControl(a, p.To, p)
|
|
||||||
if c == nil {
|
|
||||||
r.Unlock()
|
|
||||||
panic(fmt.Sprintf("No control for udp tx %s", p.To))
|
|
||||||
}
|
|
||||||
fp := r.unlockedInjectFlow(cm[x], c, p, false)
|
|
||||||
c.InjectUDPPacket(p)
|
|
||||||
fp.WasReceived()
|
|
||||||
}
|
|
||||||
r.Unlock()
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// RouteExitFunc will call the whatDo func with each udp packet from sender.
|
// RouteExitFunc will call the whatDo func with each udp packet from sender.
|
||||||
@@ -522,6 +649,7 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
switch e {
|
switch e {
|
||||||
case ExitNow:
|
case ExitNow:
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
return
|
return
|
||||||
|
|
||||||
case RouteAndExit:
|
case RouteAndExit:
|
||||||
@@ -529,6 +657,7 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
return
|
return
|
||||||
|
|
||||||
case KeepRouting:
|
case KeepRouting:
|
||||||
@@ -541,6 +670,7 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -641,6 +771,7 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
switch e {
|
switch e {
|
||||||
case ExitNow:
|
case ExitNow:
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
return
|
return
|
||||||
|
|
||||||
case RouteAndExit:
|
case RouteAndExit:
|
||||||
@@ -648,6 +779,7 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
return
|
return
|
||||||
|
|
||||||
case KeepRouting:
|
case KeepRouting:
|
||||||
@@ -659,6 +791,7 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
||||||
}
|
}
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -702,19 +835,20 @@ func (r *R) FlushAll() {
|
|||||||
}
|
}
|
||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// getControl performs or seeds NAT translation and returns the control for toAddr, p from fields may change
|
// getControl performs or seeds NAT translation and returns the control for toAddr, p from fields may change
|
||||||
// This is an internal router function, the caller must hold the lock
|
// This is an internal router function, the caller must hold the lock
|
||||||
func (r *R) getControl(fromAddr, toAddr netip.AddrPort, p *udp.Packet) *nebula.Control {
|
func (r *R) getControl(fromAddr, toAddr netip.AddrPort, p *udp.Packet) *nebula.Control {
|
||||||
if newAddr, ok := r.outNat[fromAddr.String()+":"+toAddr.String()]; ok {
|
if newAddr, ok := r.outNat[outNatKey{from: fromAddr, to: toAddr}]; ok {
|
||||||
p.From = newAddr
|
p.From = newAddr
|
||||||
}
|
}
|
||||||
|
|
||||||
c, ok := r.inNat[toAddr]
|
c, ok := r.inNat[toAddr]
|
||||||
if ok {
|
if ok {
|
||||||
r.outNat[c.GetUDPAddr().String()+":"+fromAddr.String()] = toAddr
|
r.outNat[outNatKey{from: c.GetUDPAddr(), to: fromAddr}] = toAddr
|
||||||
return c
|
return c
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,125 @@
|
|||||||
|
//go:build e2e_testing
|
||||||
|
// +build e2e_testing
|
||||||
|
|
||||||
|
package e2e
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/ed25519"
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/pem"
|
||||||
|
"net"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/cert_test"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"golang.org/x/crypto/ssh"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSSHDLifecycle(t *testing.T) {
|
||||||
|
// TestSSHDLifecycle exercises the in-process sshd through several config reloads and a Control.Stop.
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(
|
||||||
|
cert.Version1, cert.Curve_CURVE25519,
|
||||||
|
time.Now(), time.Now().Add(10*time.Minute),
|
||||||
|
nil, nil, []string{},
|
||||||
|
)
|
||||||
|
|
||||||
|
hostKeyPEM := generateSSHHostKey(t)
|
||||||
|
clientSigner, clientAuthKey := generateSSHClientKey(t)
|
||||||
|
sshdAddr := allocLoopbackPort(t)
|
||||||
|
|
||||||
|
overrides := m{
|
||||||
|
"sshd": m{
|
||||||
|
"enabled": true,
|
||||||
|
"listen": sshdAddr,
|
||||||
|
"host_key": hostKeyPEM,
|
||||||
|
"authorized_users": []m{{
|
||||||
|
"user": "tester",
|
||||||
|
"keys": []string{clientAuthKey},
|
||||||
|
}},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
control, _, _, _ := newSimpleServer(cert.Version1, ca, caKey, "sshd-test", "10.222.0.1/24", overrides)
|
||||||
|
control.Start()
|
||||||
|
t.Cleanup(func() { control.Stop() })
|
||||||
|
|
||||||
|
// sshd binds in a goroutine after Start returns; wait for it.
|
||||||
|
require.Eventually(t, func() bool { return canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
||||||
|
"sshd never started listening")
|
||||||
|
|
||||||
|
for i := 1; i <= 3; i++ {
|
||||||
|
out := sshExecReload(t, sshdAddr, clientSigner)
|
||||||
|
assert.Contains(t, out, "Reloading config", "reload cycle %d", i)
|
||||||
|
require.Eventually(t, func() bool { return canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
||||||
|
"sshd not listening after reload cycle %d", i)
|
||||||
|
}
|
||||||
|
|
||||||
|
control.Stop()
|
||||||
|
require.Eventually(t, func() bool { return !canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
||||||
|
"sshd still listening after Control.Stop")
|
||||||
|
}
|
||||||
|
|
||||||
|
func canDial(addr string) bool {
|
||||||
|
c, err := net.DialTimeout("tcp", addr, 100*time.Millisecond)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
_ = c.Close()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// allocLoopbackPort grabs an unused TCP port on 127.0.0.1, closes it, and returns the address. There
|
||||||
|
// is a small race between releasing the port and the sshd reclaiming it; in practice the OS keeps the
|
||||||
|
// port available long enough for the test to bind it.
|
||||||
|
func allocLoopbackPort(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
l, err := net.Listen("tcp", "127.0.0.1:0")
|
||||||
|
require.NoError(t, err)
|
||||||
|
addr := l.Addr().String()
|
||||||
|
require.NoError(t, l.Close())
|
||||||
|
return addr
|
||||||
|
}
|
||||||
|
|
||||||
|
func generateSSHHostKey(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
_, priv, err := ed25519.GenerateKey(rand.Reader)
|
||||||
|
require.NoError(t, err)
|
||||||
|
block, err := ssh.MarshalPrivateKey(priv, "nebula-e2e-host")
|
||||||
|
require.NoError(t, err)
|
||||||
|
return string(pem.EncodeToMemory(block))
|
||||||
|
}
|
||||||
|
|
||||||
|
func generateSSHClientKey(t *testing.T) (ssh.Signer, string) {
|
||||||
|
t.Helper()
|
||||||
|
_, priv, err := ed25519.GenerateKey(rand.Reader)
|
||||||
|
require.NoError(t, err)
|
||||||
|
signer, err := ssh.NewSignerFromKey(priv)
|
||||||
|
require.NoError(t, err)
|
||||||
|
auth := strings.TrimSpace(string(ssh.MarshalAuthorizedKey(signer.PublicKey())))
|
||||||
|
return signer, auth
|
||||||
|
}
|
||||||
|
|
||||||
|
func sshExecReload(t *testing.T, addr string, signer ssh.Signer) string {
|
||||||
|
t.Helper()
|
||||||
|
cfg := &ssh.ClientConfig{
|
||||||
|
User: "tester",
|
||||||
|
Auth: []ssh.AuthMethod{ssh.PublicKeys(signer)},
|
||||||
|
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
|
||||||
|
Timeout: 2 * time.Second,
|
||||||
|
}
|
||||||
|
client, err := ssh.Dial("tcp", addr, cfg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer client.Close()
|
||||||
|
|
||||||
|
sess, err := client.NewSession()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer sess.Close()
|
||||||
|
|
||||||
|
// reload tears the channel down before sending exit-status, so Output returns an error on the
|
||||||
|
// channel close. The output buffer still has whatever the reload callback wrote before that.
|
||||||
|
out, _ := sess.Output("reload")
|
||||||
|
return string(out)
|
||||||
|
}
|
||||||
+103
-8
@@ -15,6 +15,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
"gopkg.in/yaml.v3"
|
"gopkg.in/yaml.v3"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -42,8 +43,8 @@ func TestDropInactiveTunnels(t *testing.T) {
|
|||||||
r.Log("Go inactive and wait for the tunnels to get dropped")
|
r.Log("Go inactive and wait for the tunnels to get dropped")
|
||||||
waitStart := time.Now()
|
waitStart := time.Now()
|
||||||
for {
|
for {
|
||||||
myIndexes := len(myControl.GetHostmap().Indexes)
|
myIndexes := myControl.GetHostmapIndexCount()
|
||||||
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
theirIndexes := theirControl.GetHostmapIndexCount()
|
||||||
if myIndexes == 0 && theirIndexes == 0 {
|
if myIndexes == 0 && theirIndexes == 0 {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
@@ -355,14 +356,14 @@ func TestCrossStackRelaysWork(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80)
|
||||||
|
|
||||||
t.Log("reply?")
|
t.Log("reply?")
|
||||||
theirControl.InjectTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them")))
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnV6.Addr(), myVpnV6.Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnV6.Addr(), myVpnV6.Addr(), 80, 80)
|
||||||
|
|
||||||
@@ -373,6 +374,100 @@ func TestCrossStackRelaysWork(t *testing.T) {
|
|||||||
//relayControl.Stop()
|
//relayControl.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestRelayReplayProtection asserts that a relay (forwarding-type) node rejects
|
||||||
|
// replayed relay frames. A captured relay frame, re-injected with the same
|
||||||
|
// message counter, must be dropped by the replay window rather than re-forwarded
|
||||||
|
// to the relay target. Before the fix, handleOutsideRelayPacket authenticated the
|
||||||
|
// frame but never advanced the replay window, so every replay was re-forwarded.
|
||||||
|
func TestRelayReplayProtection(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24,fc00::1/64", m{"relay": m{"use_relays": true}})
|
||||||
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "10.128.0.128/24,fc00::128/64", m{"relay": m{"am_relay": true}})
|
||||||
|
theirUdp := netip.MustParseAddrPort("10.0.0.2:4242")
|
||||||
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServerWithUdp(cert.Version2, ca, caKey, "them ", "fc00::2/64", theirUdp, m{"relay": m{"use_relays": true}})
|
||||||
|
|
||||||
|
myVpnV6 := myVpnIpNet[1]
|
||||||
|
relayVpnV4 := relayVpnIpNet[0]
|
||||||
|
relayVpnV6 := relayVpnIpNet[1]
|
||||||
|
theirVpnV6 := theirVpnIpNet[0]
|
||||||
|
|
||||||
|
// Teach me how to reach the relay and that them is reachable via the relay
|
||||||
|
myControl.InjectLightHouseAddr(relayVpnV4.Addr(), relayUdpAddr)
|
||||||
|
myControl.InjectLightHouseAddr(relayVpnV6.Addr(), relayUdpAddr)
|
||||||
|
myControl.InjectRelays(theirVpnV6.Addr(), []netip.Addr{relayVpnV6.Addr()})
|
||||||
|
relayControl.InjectLightHouseAddr(theirVpnV6.Addr(), theirUdpAddr)
|
||||||
|
|
||||||
|
r := router.NewR(t, myControl, relayControl, theirControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
|
||||||
|
myControl.Start()
|
||||||
|
relayControl.Start()
|
||||||
|
theirControl.Start()
|
||||||
|
|
||||||
|
// Establish the relayed tunnel in both directions so all handshakes complete.
|
||||||
|
t.Log("Establish the relayed tunnel")
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me")))
|
||||||
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80)
|
||||||
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them")))
|
||||||
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
|
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnV6.Addr(), myVpnV6.Addr(), 80, 80)
|
||||||
|
|
||||||
|
// Drain anything still queued on me's UDP tx so the next packet we pull is the
|
||||||
|
// relay frame we are about to generate.
|
||||||
|
for myControl.GetFromUDP(false) != nil {
|
||||||
|
}
|
||||||
|
|
||||||
|
// Capture a single legitimate relay frame that me transmits toward the relay.
|
||||||
|
t.Log("Capture a relay frame from me -> relay")
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("replay me")))
|
||||||
|
relayFrame := myControl.GetFromUDP(true)
|
||||||
|
require.Equal(t, relayUdpAddr, relayFrame.To, "captured frame should be addressed to the relay")
|
||||||
|
var fh header.H
|
||||||
|
require.NoError(t, fh.Parse(relayFrame.Data))
|
||||||
|
require.Equal(t, header.Message, fh.Type)
|
||||||
|
require.Equal(t, header.MessageRelay, fh.Subtype)
|
||||||
|
|
||||||
|
// drainForwards counts relay frames the relay forwards toward them within the
|
||||||
|
// settle window. We match on destination + (Message, MessageRelay) so the
|
||||||
|
// relay's own direct traffic to them can't be miscounted.
|
||||||
|
drainForwards := func(settle time.Duration) int {
|
||||||
|
ch := relayControl.GetUDPTxChan()
|
||||||
|
count := 0
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case pkt := <-ch:
|
||||||
|
var ph header.H
|
||||||
|
if pkt.To == theirUdpAddr && ph.Parse(pkt.Data) == nil &&
|
||||||
|
ph.Type == header.Message && ph.Subtype == header.MessageRelay {
|
||||||
|
count++
|
||||||
|
}
|
||||||
|
pkt.Release()
|
||||||
|
case <-time.After(settle):
|
||||||
|
return count
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// First delivery of the captured frame: the relay should forward it once.
|
||||||
|
t.Log("Deliver the captured frame once; relay forwards it to them")
|
||||||
|
relayControl.InjectUDPPacket(relayFrame)
|
||||||
|
require.Equal(t, 1, drainForwards(200*time.Millisecond), "relay should forward the first, legitimate copy")
|
||||||
|
|
||||||
|
// Replay the exact same frame several times. A correct replay window rejects
|
||||||
|
// these duplicates so the relay forwards none of them.
|
||||||
|
t.Log("Replay the captured frame; relay must drop the duplicates")
|
||||||
|
const replays = 3
|
||||||
|
for i := 0; i < replays; i++ {
|
||||||
|
relayControl.InjectUDPPacket(relayFrame)
|
||||||
|
}
|
||||||
|
forwarded := drainForwards(200 * time.Millisecond)
|
||||||
|
assert.Equal(t, 0, forwarded, "relay re-forwarded %d/%d replayed relay frames; replay protection is ineffective on relay tunnels", forwarded, replays)
|
||||||
|
|
||||||
|
r.RenderHostmaps("Final hostmaps", myControl, relayControl, theirControl)
|
||||||
|
}
|
||||||
|
|
||||||
func TestCloseTunnelAuthenticated(t *testing.T) {
|
func TestCloseTunnelAuthenticated(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -398,8 +493,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) {
|
|||||||
|
|
||||||
waitStart := time.Now()
|
waitStart := time.Now()
|
||||||
for {
|
for {
|
||||||
myIndexes := len(myControl.GetHostmap().Indexes)
|
myIndexes := myControl.GetHostmapIndexCount()
|
||||||
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
theirIndexes := theirControl.GetHostmapIndexCount()
|
||||||
if myIndexes == 0 && theirIndexes == 0 {
|
if myIndexes == 0 && theirIndexes == 0 {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
@@ -453,8 +548,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) {
|
|||||||
r.Log("Injected bogus close tunnel. Let's see!")
|
r.Log("Injected bogus close tunnel. Let's see!")
|
||||||
waitStart = time.Now()
|
waitStart = time.Now()
|
||||||
for {
|
for {
|
||||||
myIndexes := len(myControl.GetHostmap().Indexes)
|
myIndexes := myControl.GetHostmapIndexCount()
|
||||||
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
theirIndexes := theirControl.GetHostmapIndexCount()
|
||||||
if myIndexes == 0 {
|
if myIndexes == 0 {
|
||||||
t.Fatal("myIndexes should not be 0")
|
t.Fatal("myIndexes should not be 0")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,188 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"log/slog"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"golang.org/x/net/ipv4"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestInnerECN(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
pkt []byte
|
||||||
|
want byte
|
||||||
|
}{
|
||||||
|
{"empty", nil, 0},
|
||||||
|
{"v4_NotECT", v4WithToS(0x00), 0x00},
|
||||||
|
{"v4_ECT0", v4WithToS(0x02), 0x02},
|
||||||
|
{"v4_ECT1", v4WithToS(0x01), 0x01},
|
||||||
|
{"v4_CE", v4WithToS(0x03), 0x03},
|
||||||
|
{"v4_DSCP_then_NotECT", v4WithToS(0x88 | 0x00), 0x00},
|
||||||
|
{"v4_DSCP_then_CE", v4WithToS(0x88 | 0x03), 0x03},
|
||||||
|
{"v6_NotECT", v6WithTC(0x00), 0x00},
|
||||||
|
{"v6_ECT0", v6WithTC(0x02), 0x02},
|
||||||
|
{"v6_CE", v6WithTC(0x03), 0x03},
|
||||||
|
{"v6_DSCP_then_CE", v6WithTC(0x88 | 0x03), 0x03},
|
||||||
|
{"unknown_version", []byte{0xa5, 0xff}, 0},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
t.Run(c.name, func(t *testing.T) {
|
||||||
|
got := innerECN(c.pkt)
|
||||||
|
if got != c.want {
|
||||||
|
t.Errorf("innerECN=0x%02x want 0x%02x", got, c.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// v4WithToS returns a 2-byte slice tall enough for innerECN: byte 0 carries
|
||||||
|
// version=4 in the high nibble, byte 1 is the full ToS so we exercise both
|
||||||
|
// the DSCP and ECN portions through the byte 1 mask.
|
||||||
|
func v4WithToS(tos byte) []byte {
|
||||||
|
return []byte{0x45, tos}
|
||||||
|
}
|
||||||
|
|
||||||
|
// v6WithTC builds a 2-byte slice that places a known traffic class value
|
||||||
|
// across bytes 0 (high nibble of TC) and 1 (low nibble of TC). innerECN
|
||||||
|
// extracts ECN as (b[1]>>4)&0x03, which corresponds to TC[1:0].
|
||||||
|
func v6WithTC(tc byte) []byte {
|
||||||
|
return []byte{0x60 | (tc>>4)&0x0f, (tc & 0x0f) << 4}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestApplyOuterECN(t *testing.T) {
|
||||||
|
silent := slog.New(slog.DiscardHandler)
|
||||||
|
hi := &HostInfo{}
|
||||||
|
|
||||||
|
// Build a v4 packet helper with a given inner ECN field.
|
||||||
|
v4 := func(innerECN byte) []byte {
|
||||||
|
// 20-byte minimal IPv4 header with ToS = innerECN (DSCP zeroed).
|
||||||
|
return []byte{
|
||||||
|
0x45, innerECN, 0, 28,
|
||||||
|
0, 0, 0x40, 0,
|
||||||
|
64, 6, 0, 0,
|
||||||
|
10, 0, 0, 1,
|
||||||
|
10, 0, 0, 2,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Build a v6 packet helper with a given inner ECN field. ECN occupies
|
||||||
|
// TC[1:0] which sit at byte 1 mask 0x30.
|
||||||
|
v6 := func(innerECN byte) []byte {
|
||||||
|
// 40-byte minimal IPv6 header with TC[1:0] = innerECN.
|
||||||
|
pkt := make([]byte, 40)
|
||||||
|
pkt[0] = 0x60 // version=6, TC[7:4]=0
|
||||||
|
pkt[1] = (innerECN & 0x03) << 4 // TC[3:0]: low 2 bits = ECN, top 2 = DSCP-low (0)
|
||||||
|
return pkt
|
||||||
|
}
|
||||||
|
|
||||||
|
type cell struct {
|
||||||
|
outer byte
|
||||||
|
inner byte
|
||||||
|
wantECN byte
|
||||||
|
wantSame bool // expect inner unchanged (true => verify the byte didn't move)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RFC 6040 normal-mode combine table. Only outer==CE causes mutation.
|
||||||
|
table := []cell{
|
||||||
|
{ecnNotECT, ecnNotECT, ecnNotECT, true},
|
||||||
|
{ecnNotECT, ecnECT0, ecnECT0, true},
|
||||||
|
{ecnNotECT, ecnECT1, ecnECT1, true},
|
||||||
|
{ecnNotECT, ecnCE, ecnCE, true},
|
||||||
|
|
||||||
|
{ecnECT0, ecnNotECT, ecnNotECT, true},
|
||||||
|
{ecnECT0, ecnECT0, ecnECT0, true},
|
||||||
|
{ecnECT0, ecnECT1, ecnECT1, true},
|
||||||
|
{ecnECT0, ecnCE, ecnCE, true},
|
||||||
|
|
||||||
|
{ecnECT1, ecnNotECT, ecnNotECT, true},
|
||||||
|
{ecnECT1, ecnECT0, ecnECT0, true},
|
||||||
|
{ecnECT1, ecnECT1, ecnECT1, true},
|
||||||
|
{ecnECT1, ecnCE, ecnCE, true},
|
||||||
|
|
||||||
|
{ecnCE, ecnNotECT, ecnNotECT, true}, // legacy: log, leave alone
|
||||||
|
{ecnCE, ecnECT0, ecnCE, false}, // CE folded in
|
||||||
|
{ecnCE, ecnECT1, ecnCE, false},
|
||||||
|
{ecnCE, ecnCE, ecnCE, true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, c := range table {
|
||||||
|
t.Run("v4", func(t *testing.T) {
|
||||||
|
pkt := v4(c.inner)
|
||||||
|
applyOuterECN(pkt, c.outer, hi, silent)
|
||||||
|
got := pkt[1] & 0x03
|
||||||
|
if got != c.wantECN {
|
||||||
|
t.Errorf("v4 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
t.Run("v6", func(t *testing.T) {
|
||||||
|
pkt := v6(c.inner)
|
||||||
|
applyOuterECN(pkt, c.outer, hi, silent)
|
||||||
|
got := (pkt[1] >> 4) & 0x03
|
||||||
|
if got != c.wantECN {
|
||||||
|
t.Errorf("v6 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestApplyOuterECN_IPv4ChecksumStaysValid guards against H1: folding an outer
|
||||||
|
// CE mark into the inner IPv4 ToS byte must keep the IPv4 header checksum valid.
|
||||||
|
// The passthrough emit paths write the packet verbatim, so a stale checksum
|
||||||
|
// turns an underlay congestion mark into packet loss at the receiver.
|
||||||
|
func TestApplyOuterECN_IPv4ChecksumStaysValid(t *testing.T) {
|
||||||
|
silent := slog.New(slog.DiscardHandler)
|
||||||
|
hi := &HostInfo{}
|
||||||
|
|
||||||
|
// 20-byte IPv4 header with DSCP=0x88 and inner ECN = ECT(0). Folding CE
|
||||||
|
// flips only the low two bits of the ToS byte while leaving DSCP intact.
|
||||||
|
pkt := []byte{
|
||||||
|
0x45, 0x88 | ecnECT0, 0, 40,
|
||||||
|
0x1c, 0x46, 0x40, 0x00,
|
||||||
|
64, 6, 0, 0,
|
||||||
|
10, 0, 0, 1,
|
||||||
|
10, 0, 0, 2,
|
||||||
|
}
|
||||||
|
// Stamp a correct header checksum before the fold.
|
||||||
|
binary.BigEndian.PutUint16(pkt[10:12], ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]))
|
||||||
|
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
|
||||||
|
t.Fatal("test setup: initial header checksum invalid")
|
||||||
|
}
|
||||||
|
|
||||||
|
applyOuterECN(pkt, ecnCE, hi, silent)
|
||||||
|
|
||||||
|
// CE folded in, DSCP preserved.
|
||||||
|
if got, want := pkt[1], byte(0x88|ecnCE); got != want {
|
||||||
|
t.Fatalf("ToS after fold = 0x%02x, want 0x%02x", got, want)
|
||||||
|
}
|
||||||
|
// The incremental RFC 1624 update must leave the checksum valid and equal
|
||||||
|
// to a full recompute over the mutated header.
|
||||||
|
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
|
||||||
|
t.Fatalf("IPv4 header checksum invalid after CE fold: 0x%04x", binary.BigEndian.Uint16(pkt[10:12]))
|
||||||
|
}
|
||||||
|
if got, want := binary.BigEndian.Uint16(pkt[10:12]), ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]); got != want {
|
||||||
|
t.Fatalf("checksum = 0x%04x, full recompute = 0x%04x", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ipv4HeaderChecksum computes the RFC 1071 IPv4 header checksum over hdr,
|
||||||
|
// treating the checksum field (bytes 10:12) as zero.
|
||||||
|
func ipv4HeaderChecksum(hdr []byte) uint16 {
|
||||||
|
var sum uint32
|
||||||
|
for i := 0; i+1 < len(hdr); i += 2 {
|
||||||
|
if i == 10 {
|
||||||
|
continue // checksum field
|
||||||
|
}
|
||||||
|
sum += uint32(hdr[i])<<8 | uint32(hdr[i+1])
|
||||||
|
}
|
||||||
|
for sum > 0xffff {
|
||||||
|
sum = (sum >> 16) + (sum & 0xffff)
|
||||||
|
}
|
||||||
|
return ^uint16(sum)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ipv4HeaderChecksumValid reports whether the stored checksum matches a fresh
|
||||||
|
// computation over the header.
|
||||||
|
func ipv4HeaderChecksumValid(hdr []byte) bool {
|
||||||
|
return binary.BigEndian.Uint16(hdr[10:12]) == ipv4HeaderChecksum(hdr)
|
||||||
|
}
|
||||||
+63
-1
@@ -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
|
||||||
@@ -242,6 +254,28 @@ tun:
|
|||||||
# Default MTU for every packet, safe setting is (and the default) 1300 for internet based traffic
|
# Default MTU for every packet, safe setting is (and the default) 1300 for internet based traffic
|
||||||
mtu: 1300
|
mtu: 1300
|
||||||
|
|
||||||
|
# Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's
|
||||||
|
# batched sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire
|
||||||
|
# instead of being sprayed across multiple TX rings and reordered. Not reloadable.
|
||||||
|
#
|
||||||
|
# When cpu_affinity is unset, nebula picks CPUs that do NOT service any physical NIC's interrupts (read from
|
||||||
|
# /sys/class/net/*/device/msi_irqs and /proc/irq/*/effective_affinity_list): an encrypt thread pinned onto a core
|
||||||
|
# that also runs NAPI for a NIC RX queue fights the softirq for the core and collapses throughput for flows hashed
|
||||||
|
# to that queue. If the NIC's vectors blanket every allowed CPU (many drivers default to one queue per core) the
|
||||||
|
# avoidance logs and falls back to the old spread; narrow the NIC's queue/IRQ spread (e.g. `ethtool -X <dev>
|
||||||
|
# equal N`) or set cpu_affinity explicitly to benefit.
|
||||||
|
#pin_threads: true
|
||||||
|
|
||||||
|
# Linux only. cpu_affinity overrides which CPUs the tun reader threads pin to: a list of CPU IDs, one per routine
|
||||||
|
# (see the top-level `routines` setting). Lists shorter than `routines` are modulo-cycled across the queues; extra
|
||||||
|
# entries are ignored. IDs must be within the process's allowed CPU set, so this respects taskset / cgroup cpusets;
|
||||||
|
# a non-integer or not-allowed entry disables the override and falls back to spreading queues across the allowed
|
||||||
|
# CPUs. Setting this disables the automatic NIC-IRQ avoidance described under pin_threads — prefer CPUs that don't
|
||||||
|
# service your underlay NIC's RX queue IRQs. Only meaningful while pin_threads is true. Not reloadable.
|
||||||
|
#cpu_affinity:
|
||||||
|
# - 2
|
||||||
|
# - 4
|
||||||
|
|
||||||
# Route based MTU overrides, you have known vpn ip paths that can support larger MTUs you can increase/decrease them here
|
# Route based MTU overrides, you have known vpn ip paths that can support larger MTUs you can increase/decrease them here
|
||||||
routes:
|
routes:
|
||||||
#- mtu: 8800
|
#- mtu: 8800
|
||||||
@@ -282,6 +316,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
|
||||||
@@ -360,6 +412,16 @@ logging:
|
|||||||
# This setting is reloadable
|
# This setting is reloadable
|
||||||
#inactivity_timeout: 10m
|
#inactivity_timeout: 10m
|
||||||
|
|
||||||
|
# ecn (default true) propagates ECN (Explicit Congestion Notification) across the tunnel per RFC 6040: the inner
|
||||||
|
# packet's ECN codepoint is copied onto the outer carrier header on encapsulation, and an outer CE ("congestion
|
||||||
|
# experienced") mark is folded back into the inner header on decapsulation. On linux it additionally stamps
|
||||||
|
# RTAX_FEATURE_ECN on the routes nebula installs, so the kernel actively negotiates ECN for connections to mesh
|
||||||
|
# prefixes. Disable this only when an underlay middlebox mangles or clears ECN bits unpredictably.
|
||||||
|
# This setting is reloadable, BUT flipping it at runtime only updates the datapath (the inner<->outer copy/combine).
|
||||||
|
# The RTAX_FEATURE_ECN flag on already-installed routes is NOT revisited on reload, so nebula must be restarted for
|
||||||
|
# the route half of this setting to take effect.
|
||||||
|
#ecn: true
|
||||||
|
|
||||||
# Nebula security group configuration
|
# Nebula security group configuration
|
||||||
firewall:
|
firewall:
|
||||||
# Action to take when a packet is not allowed by the firewall rules.
|
# Action to take when a packet is not allowed by the firewall rules.
|
||||||
@@ -367,7 +429,7 @@ firewall:
|
|||||||
# `drop` (default): silently drop the packet.
|
# `drop` (default): silently drop the packet.
|
||||||
# `reject`: send a reject reply.
|
# `reject`: send a reject reply.
|
||||||
# - For TCP, this will be a RST "Connection Reset" packet.
|
# - For TCP, this will be a RST "Connection Reset" packet.
|
||||||
# - For other protocols, this will be an ICMP port unreachable packet.
|
# - For other protocols, this will be an ICMP "Destination unreachable: Communication administratively prohibited" packet.
|
||||||
outbound_action: drop
|
outbound_action: drop
|
||||||
inbound_action: drop
|
inbound_action: drop
|
||||||
|
|
||||||
|
|||||||
+56
-39
@@ -44,8 +44,8 @@ type Firewall struct {
|
|||||||
InRules *FirewallTable
|
InRules *FirewallTable
|
||||||
OutRules *FirewallTable
|
OutRules *FirewallTable
|
||||||
|
|
||||||
InSendReject bool
|
InboundSendReject bool
|
||||||
OutSendReject bool
|
OutboundSendReject bool
|
||||||
|
|
||||||
//TODO: we should have many more options for TCP, an option for ICMP, and mimic the kernel a bit better
|
//TODO: we should have many more options for TCP, an option for ICMP, and mimic the kernel a bit better
|
||||||
// https://www.kernel.org/doc/Documentation/networking/nf_conntrack-sysctl.txt
|
// https://www.kernel.org/doc/Documentation/networking/nf_conntrack-sysctl.txt
|
||||||
@@ -58,8 +58,9 @@ type Firewall struct {
|
|||||||
routableNetworks *bart.Lite
|
routableNetworks *bart.Lite
|
||||||
|
|
||||||
// assignedNetworks is a list of vpn networks assigned to us in the certificate.
|
// assignedNetworks is a list of vpn networks assigned to us in the certificate.
|
||||||
assignedNetworks []netip.Prefix
|
assignedNetworks []netip.Prefix
|
||||||
hasUnsafeNetworks bool
|
// unsafeNetworks is the list of unsafe networks issued to us in the certificate
|
||||||
|
unsafeNetworks []netip.Prefix
|
||||||
|
|
||||||
rules string
|
rules string
|
||||||
rulesVersion uint16
|
rulesVersion uint16
|
||||||
@@ -158,10 +159,9 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
|
|||||||
assignedNetworks = append(assignedNetworks, network)
|
assignedNetworks = append(assignedNetworks, network)
|
||||||
}
|
}
|
||||||
|
|
||||||
hasUnsafeNetworks := false
|
unsafeNetworks := c.UnsafeNetworks()
|
||||||
for _, n := range c.UnsafeNetworks() {
|
for _, n := range unsafeNetworks {
|
||||||
routableNetworks.Insert(n)
|
routableNetworks.Insert(n)
|
||||||
hasUnsafeNetworks = true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return &Firewall{
|
return &Firewall{
|
||||||
@@ -169,15 +169,15 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
|
|||||||
Conns: make(map[firewall.Packet]*conn),
|
Conns: make(map[firewall.Packet]*conn),
|
||||||
TimerWheel: NewTimerWheel[firewall.Packet](tmin, tmax),
|
TimerWheel: NewTimerWheel[firewall.Packet](tmin, tmax),
|
||||||
},
|
},
|
||||||
InRules: newFirewallTable(),
|
InRules: newFirewallTable(),
|
||||||
OutRules: newFirewallTable(),
|
OutRules: newFirewallTable(),
|
||||||
TCPTimeout: tcpTimeout,
|
TCPTimeout: tcpTimeout,
|
||||||
UDPTimeout: UDPTimeout,
|
UDPTimeout: UDPTimeout,
|
||||||
DefaultTimeout: defaultTimeout,
|
DefaultTimeout: defaultTimeout,
|
||||||
routableNetworks: routableNetworks,
|
routableNetworks: routableNetworks,
|
||||||
assignedNetworks: assignedNetworks,
|
assignedNetworks: assignedNetworks,
|
||||||
hasUnsafeNetworks: hasUnsafeNetworks,
|
unsafeNetworks: unsafeNetworks,
|
||||||
l: l,
|
l: l,
|
||||||
|
|
||||||
incomingMetrics: firewallMetrics{
|
incomingMetrics: firewallMetrics{
|
||||||
droppedLocalAddr: metrics.GetOrRegisterCounter("firewall.incoming.dropped.local_addr", nil),
|
droppedLocalAddr: metrics.GetOrRegisterCounter("firewall.incoming.dropped.local_addr", nil),
|
||||||
@@ -216,23 +216,23 @@ func NewFirewallFromConfig(l *slog.Logger, cs *CertState, c *config.C) (*Firewal
|
|||||||
inboundAction := c.GetString("firewall.inbound_action", "drop")
|
inboundAction := c.GetString("firewall.inbound_action", "drop")
|
||||||
switch inboundAction {
|
switch inboundAction {
|
||||||
case "reject":
|
case "reject":
|
||||||
fw.InSendReject = true
|
fw.InboundSendReject = true
|
||||||
case "drop":
|
case "drop":
|
||||||
fw.InSendReject = false
|
fw.InboundSendReject = false
|
||||||
default:
|
default:
|
||||||
l.Warn("invalid firewall.inbound_action, defaulting to `drop`", "action", inboundAction)
|
l.Warn("invalid firewall.inbound_action, defaulting to `drop`", "action", inboundAction)
|
||||||
fw.InSendReject = false
|
fw.InboundSendReject = false
|
||||||
}
|
}
|
||||||
|
|
||||||
outboundAction := c.GetString("firewall.outbound_action", "drop")
|
outboundAction := c.GetString("firewall.outbound_action", "drop")
|
||||||
switch outboundAction {
|
switch outboundAction {
|
||||||
case "reject":
|
case "reject":
|
||||||
fw.OutSendReject = true
|
fw.OutboundSendReject = true
|
||||||
case "drop":
|
case "drop":
|
||||||
fw.OutSendReject = false
|
fw.OutboundSendReject = false
|
||||||
default:
|
default:
|
||||||
l.Warn("invalid firewall.outbound_action, defaulting to `drop`", "action", outboundAction)
|
l.Warn("invalid firewall.outbound_action, defaulting to `drop`", "action", outboundAction)
|
||||||
fw.OutSendReject = false
|
fw.OutboundSendReject = false
|
||||||
}
|
}
|
||||||
|
|
||||||
err := AddFirewallRulesFromConfig(l, false, c, fw)
|
err := AddFirewallRulesFromConfig(l, false, c, fw)
|
||||||
@@ -423,11 +423,6 @@ var ErrNoMatchingRule = errors.New("no matching rule in firewall table")
|
|||||||
// Drop returns an error if the packet should be dropped, explaining why. It
|
// Drop returns an error if the packet should be dropped, explaining why. It
|
||||||
// returns nil if the packet should not be dropped.
|
// returns nil if the packet should not be dropped.
|
||||||
func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) error {
|
func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) error {
|
||||||
// Check if we spoke to this tuple, if we did then allow this packet
|
|
||||||
if f.inConns(fp, h, caPool, localCache) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Make sure remote address matches nebula certificate, and determine how to treat it
|
// Make sure remote address matches nebula certificate, and determine how to treat it
|
||||||
if h.networks == nil {
|
if h.networks == nil {
|
||||||
// Simple case: Certificate has one address and no unsafe networks
|
// Simple case: Certificate has one address and no unsafe networks
|
||||||
@@ -461,6 +456,11 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *
|
|||||||
return ErrInvalidLocalIP
|
return ErrInvalidLocalIP
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Check if we spoke to this tuple, if we did then allow this packet
|
||||||
|
if f.inConns(fp, h, caPool, localCache) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
table := f.OutRules
|
table := f.OutRules
|
||||||
if incoming {
|
if incoming {
|
||||||
table = f.InRules
|
table = f.InRules
|
||||||
@@ -897,7 +897,7 @@ func (flc *firewallLocalCIDR) addRule(f *Firewall, localCidr string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if localCidr == "" {
|
if localCidr == "" {
|
||||||
if !f.hasUnsafeNetworks || f.defaultLocalCIDRAny {
|
if len(f.unsafeNetworks) == 0 || f.defaultLocalCIDRAny {
|
||||||
flc.Any = true
|
flc.Any = true
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -1055,7 +1055,6 @@ func (r *rule) sanity() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func parsePort(s string) (int32, int32, error) {
|
func parsePort(s string) (int32, int32, error) {
|
||||||
var err error
|
|
||||||
const notAPort int32 = -2
|
const notAPort int32 = -2
|
||||||
if s == "any" {
|
if s == "any" {
|
||||||
return firewall.PortAny, firewall.PortAny, nil
|
return firewall.PortAny, firewall.PortAny, nil
|
||||||
@@ -1064,11 +1063,11 @@ func parsePort(s string) (int32, int32, error) {
|
|||||||
return firewall.PortFragment, firewall.PortFragment, nil
|
return firewall.PortFragment, firewall.PortFragment, nil
|
||||||
}
|
}
|
||||||
if !strings.Contains(s, `-`) {
|
if !strings.Contains(s, `-`) {
|
||||||
rPort, err := strconv.Atoi(s)
|
rPort, err := parsePortValue("", s)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return notAPort, notAPort, fmt.Errorf("was not a number; `%s`", s)
|
return notAPort, notAPort, err
|
||||||
}
|
}
|
||||||
return int32(rPort), int32(rPort), nil
|
return rPort, rPort, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
sPorts := strings.SplitN(s, `-`, 2)
|
sPorts := strings.SplitN(s, `-`, 2)
|
||||||
@@ -1079,22 +1078,40 @@ func parsePort(s string) (int32, int32, error) {
|
|||||||
return notAPort, notAPort, fmt.Errorf("appears to be a range but could not be parsed; `%s`", s)
|
return notAPort, notAPort, fmt.Errorf("appears to be a range but could not be parsed; `%s`", s)
|
||||||
}
|
}
|
||||||
|
|
||||||
rStartPort, err := strconv.Atoi(sPorts[0])
|
startPort, err := parsePortValue("beginning range ", sPorts[0])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return notAPort, notAPort, fmt.Errorf("beginning range was not a number; `%s`", sPorts[0])
|
return notAPort, notAPort, err
|
||||||
}
|
}
|
||||||
|
|
||||||
rEndPort, err := strconv.Atoi(sPorts[1])
|
endPort, err := parsePortValue("ending range ", sPorts[1])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return notAPort, notAPort, fmt.Errorf("ending range was not a number; `%s`", sPorts[1])
|
return notAPort, notAPort, err
|
||||||
}
|
}
|
||||||
|
|
||||||
startPort := int32(rStartPort)
|
|
||||||
endPort := int32(rEndPort)
|
|
||||||
|
|
||||||
if startPort == firewall.PortAny {
|
if startPort == firewall.PortAny {
|
||||||
endPort = firewall.PortAny
|
endPort = firewall.PortAny
|
||||||
}
|
}
|
||||||
|
|
||||||
return startPort, endPort, nil
|
return startPort, endPort, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// parsePortValue accepts a base-10 decimal in [0, 65535] and returns it
|
||||||
|
// widened to int32. Using strconv.ParseUint with bitSize 16 rejects
|
||||||
|
// negative input, out-of-range input (>65535), and any non-decimal byte
|
||||||
|
// by construction, so the int32 widening that follows is provably safe
|
||||||
|
// and cannot collide with firewall.PortAny (0) or firewall.PortFragment
|
||||||
|
// (-1) via integer truncation.
|
||||||
|
//
|
||||||
|
// prefix is prepended to both error messages so callers can disambiguate
|
||||||
|
// the single-port path (prefix="") from the range bounds (prefix="beginning
|
||||||
|
// range " / "ending range "), preserving the historical error strings.
|
||||||
|
func parsePortValue(prefix, s string) (int32, error) {
|
||||||
|
n, err := strconv.ParseUint(s, 10, 16)
|
||||||
|
if err == nil {
|
||||||
|
return int32(n), nil
|
||||||
|
}
|
||||||
|
if errors.Is(err, strconv.ErrRange) {
|
||||||
|
return 0, fmt.Errorf("%sout of range [0,65535]; `%s`", prefix, s)
|
||||||
|
}
|
||||||
|
return 0, fmt.Errorf("%swas not a number; `%s`", prefix, s)
|
||||||
|
}
|
||||||
|
|||||||
+4
-2
@@ -5,6 +5,8 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/logging"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ConntrackCache is used as a local routine cache to know if a given flow
|
// ConntrackCache is used as a local routine cache to know if a given flow
|
||||||
@@ -56,8 +58,8 @@ func (c *ConntrackCacheTicker) Get() ConntrackCache {
|
|||||||
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
||||||
c.cacheV = tick
|
c.cacheV = tick
|
||||||
if ll := len(c.cache); ll > 0 {
|
if ll := len(c.cache); ll > 0 {
|
||||||
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
if c.l.Enabled(context.Background(), logging.LevelTrace) {
|
||||||
c.l.Debug("resetting conntrack cache", "len", ll)
|
c.l.Log(context.Background(), logging.LevelTrace, "resetting conntrack cache", "len", ll)
|
||||||
}
|
}
|
||||||
c.cache = make(ConntrackCache, ll)
|
c.cache = make(ConntrackCache, ll)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/logging"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
@@ -30,27 +31,27 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
|
|||||||
|
|
||||||
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 3)
|
c := newFixedTicker(t, l, 3)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|
||||||
assert.Equal(t, "level=DEBUG msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
assert.Equal(t, "level=DEBUG-4 msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewJSONLoggerWithOutput(buf, slog.LevelDebug)
|
l := test.NewJSONLoggerWithOutput(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 2)
|
c := newFixedTicker(t, l, 2)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|
||||||
assert.JSONEq(t, `{"level":"DEBUG","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
assert.JSONEq(t, `{"level":"DEBUG-4","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelInfo)
|
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 5)
|
c := newFixedTicker(t, l, 5)
|
||||||
c.Get()
|
c.Get()
|
||||||
@@ -60,7 +61,7 @@ func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
|||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 0)
|
c := newFixedTicker(t, l, 0)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|||||||
@@ -916,6 +916,159 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
|
|||||||
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
|
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestFirewall_ConntrackSourceSpoofingAcrossPeers(t *testing.T) {
|
||||||
|
l := test.NewLoggerWithOutput(&bytes.Buffer{})
|
||||||
|
|
||||||
|
myVpnNetworksTable := new(bart.Lite)
|
||||||
|
myVpnNetworksTable.Insert(netip.MustParsePrefix("192.0.2.1/24"))
|
||||||
|
|
||||||
|
owner := &dummyCert{
|
||||||
|
name: "owner",
|
||||||
|
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.1/24")},
|
||||||
|
}
|
||||||
|
|
||||||
|
victim := &cert.CachedCertificate{
|
||||||
|
Certificate: &dummyCert{
|
||||||
|
name: "victim",
|
||||||
|
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
victimHI := HostInfo{
|
||||||
|
ConnectionState: &ConnectionState{peerCert: victim},
|
||||||
|
vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")},
|
||||||
|
}
|
||||||
|
victimHI.buildNetworks(myVpnNetworksTable, victim.Certificate)
|
||||||
|
|
||||||
|
attacker := &cert.CachedCertificate{
|
||||||
|
Certificate: &dummyCert{
|
||||||
|
name: "attacker",
|
||||||
|
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.3/24")},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
attackerHI := HostInfo{
|
||||||
|
ConnectionState: &ConnectionState{peerCert: attacker},
|
||||||
|
vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.3")},
|
||||||
|
}
|
||||||
|
attackerHI.buildNetworks(myVpnNetworksTable, attacker.Certificate)
|
||||||
|
|
||||||
|
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, owner)
|
||||||
|
// Allow any inbound traffic that passes the cert / source-IP checks.
|
||||||
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", ""))
|
||||||
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
|
flow := firewall.Packet{
|
||||||
|
LocalAddr: netip.MustParseAddr("192.0.2.1"),
|
||||||
|
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
||||||
|
LocalPort: 443,
|
||||||
|
RemotePort: 55000,
|
||||||
|
Protocol: firewall.ProtoUDP,
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil),
|
||||||
|
"victim's own traffic from its own overlay IP must be allowed")
|
||||||
|
|
||||||
|
unseen := flow
|
||||||
|
unseen.RemotePort = 55001
|
||||||
|
assert.Equal(t, ErrInvalidRemoteIP, fw.Drop(unseen, true, &attackerHI, cp, nil),
|
||||||
|
"sanity: attacker forging victim's source IP must be rejected when no conntrack entry exists")
|
||||||
|
|
||||||
|
got := fw.Drop(flow, true, &attackerHI, cp, nil)
|
||||||
|
t.Logf("attacker replaying victim's 4-tuple: Drop returned %v (nil == packet ALLOWED == spoof succeeded)", got)
|
||||||
|
assert.Equal(t, ErrInvalidRemoteIP, got,
|
||||||
|
"SECURITY: attacker spoofed victim's overlay source IP (192.0.2.2) by reusing an existing conntrack 4-tuple; Drop returned %v instead of rejecting", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkFirewallDropConntrackHit measures Drop on an already-established flow
|
||||||
|
// (a conntrack hit). This is the fast path that the source-IP<->cert binding
|
||||||
|
// reordering adds work to, so it quantifies the cost of moving the address checks
|
||||||
|
// ahead of the conntrack lookup. Cases:
|
||||||
|
// - simple: peer cert has one address, no unsafe networks (h.networks == nil),
|
||||||
|
// so the remote-address check is a single netip.Addr compare.
|
||||||
|
// - complex: peer cert has unsafe networks (h.networks populated), so the
|
||||||
|
// remote-address check is a BART lookup.
|
||||||
|
// - noCache/localCache: whether a per-batch ConntrackCache is supplied, which in
|
||||||
|
// the original code let the fast path skip straight past the address checks.
|
||||||
|
func BenchmarkFirewallDropConntrackHit(b *testing.B) {
|
||||||
|
l := test.NewLoggerWithOutput(&bytes.Buffer{})
|
||||||
|
|
||||||
|
myVpnNetworksTable := new(bart.Lite)
|
||||||
|
myVpnNetworksTable.Insert(netip.MustParsePrefix("192.0.2.1/24"))
|
||||||
|
|
||||||
|
owner := &dummyCert{
|
||||||
|
name: "owner",
|
||||||
|
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.1/24")},
|
||||||
|
}
|
||||||
|
|
||||||
|
simpleCert := &cert.CachedCertificate{
|
||||||
|
Certificate: &dummyCert{
|
||||||
|
name: "simple",
|
||||||
|
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
simpleHost := &HostInfo{
|
||||||
|
ConnectionState: &ConnectionState{peerCert: simpleCert},
|
||||||
|
vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")},
|
||||||
|
}
|
||||||
|
simpleHost.buildNetworks(myVpnNetworksTable, simpleCert.Certificate)
|
||||||
|
|
||||||
|
complexCert := &cert.CachedCertificate{
|
||||||
|
Certificate: &dummyCert{
|
||||||
|
name: "complex",
|
||||||
|
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")},
|
||||||
|
unsafeNetworks: []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
complexHost := &HostInfo{
|
||||||
|
ConnectionState: &ConnectionState{peerCert: complexCert},
|
||||||
|
vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")},
|
||||||
|
}
|
||||||
|
complexHost.buildNetworks(myVpnNetworksTable, complexCert.Certificate)
|
||||||
|
|
||||||
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
|
flow := firewall.Packet{
|
||||||
|
LocalAddr: netip.MustParseAddr("192.0.2.1"),
|
||||||
|
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
||||||
|
LocalPort: 443,
|
||||||
|
RemotePort: 55000,
|
||||||
|
Protocol: firewall.ProtoUDP,
|
||||||
|
}
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
host *HostInfo
|
||||||
|
useCache bool
|
||||||
|
}{
|
||||||
|
{"simple/noCache", simpleHost, false},
|
||||||
|
{"simple/localCache", simpleHost, true},
|
||||||
|
{"complex/noCache", complexHost, false},
|
||||||
|
{"complex/localCache", complexHost, true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range cases {
|
||||||
|
b.Run(tc.name, func(b *testing.B) {
|
||||||
|
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, owner)
|
||||||
|
require.NoError(b, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", ""))
|
||||||
|
|
||||||
|
// Establish the conntrack entry so every benchmarked Drop is a hit.
|
||||||
|
require.NoError(b, fw.Drop(flow, true, tc.host, cp, nil))
|
||||||
|
|
||||||
|
var cache firewall.ConntrackCache
|
||||||
|
if tc.useCache {
|
||||||
|
cache = firewall.ConntrackCache{}
|
||||||
|
}
|
||||||
|
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
if err := fw.Drop(flow, true, tc.host, cp, cache); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func BenchmarkLookup(b *testing.B) {
|
func BenchmarkLookup(b *testing.B) {
|
||||||
ml := func(m map[string]struct{}, a [][]string) {
|
ml := func(m map[string]struct{}, a [][]string) {
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
@@ -1029,6 +1182,75 @@ func Test_parsePort(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Test_parsePort_invalid covers inputs that must error. The named bug is
|
||||||
|
// that int32(strconv.Atoi("4294967296")) truncates to 0 == firewall.PortAny,
|
||||||
|
// silently turning a typo into a match-all-ports rule; the rest are
|
||||||
|
// representative syntax/range probes.
|
||||||
|
func Test_parsePort_invalid(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
wantErrContains string
|
||||||
|
}{
|
||||||
|
// Numeric overflow (the named bug + boundary).
|
||||||
|
{"named bug: 2^32 truncates to PortAny", "4294967296", "out of range"},
|
||||||
|
{"just above max real port", "65536", "out of range"},
|
||||||
|
|
||||||
|
// Negatives route through the range branch and hit the empty-half
|
||||||
|
// guard; included as defense in depth so a future refactor cannot
|
||||||
|
// accidentally reach the int32 cast.
|
||||||
|
{"negative", "-1", "could not be parsed"},
|
||||||
|
|
||||||
|
// Syntax probes.
|
||||||
|
{"NUL between digits", "4\x002", "was not a number"},
|
||||||
|
{"hex notation", "0x10", "was not a number"},
|
||||||
|
{"scientific notation", "1e3", "was not a number"},
|
||||||
|
{"leading whitespace", " 42", "was not a number"},
|
||||||
|
{"fullwidth digits", "42", "was not a number"},
|
||||||
|
|
||||||
|
// Range branch.
|
||||||
|
{"range upper out of range", "1-65536", "ending range out of range"},
|
||||||
|
{"range lower out of range", "65536-65537", "beginning range out of range"},
|
||||||
|
{"range with negative upper", "1--1", "ending range was not a number"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
_, _, err := parsePort(tc.input)
|
||||||
|
require.Error(t, err, "input %q must error", tc.input)
|
||||||
|
require.ErrorContains(t, err, tc.wantErrContains)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test_parsePort_valid_boundaries locks in success cases at 0, 1, and 65535
|
||||||
|
// so a future refactor cannot regress the boundaries.
|
||||||
|
func Test_parsePort_valid_boundaries(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
wantStart int32
|
||||||
|
wantEnd int32
|
||||||
|
}{
|
||||||
|
{"zero is PortAny", "0", 0, 0},
|
||||||
|
{"min real port", "1", 1, 1},
|
||||||
|
{"max real port", "65535", 65535, 65535},
|
||||||
|
{"range zero to max forces end to zero", "0-65535", 0, 0},
|
||||||
|
{"range max to max", "65535-65535", 65535, 65535},
|
||||||
|
{"range one to max", "1-65535", 1, 65535},
|
||||||
|
{"range with whitespace inside", " 1 - 2 ", 1, 2},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
s, e, err := parsePort(tc.input)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, tc.wantStart, s, "start port")
|
||||||
|
assert.Equal(t, tc.wantEnd, e, "end port")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestNewFirewallFromConfig(t *testing.T) {
|
func TestNewFirewallFromConfig(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
// Test a bad rule definition
|
// Test a bad rule definition
|
||||||
|
|||||||
@@ -9,10 +9,10 @@ require (
|
|||||||
github.com/armon/go-radix v1.0.0
|
github.com/armon/go-radix v1.0.0
|
||||||
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432
|
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432
|
||||||
github.com/flynn/noise v1.1.0
|
github.com/flynn/noise v1.1.0
|
||||||
github.com/gaissmai/bart v0.26.0
|
github.com/gaissmai/bart v0.28.0
|
||||||
github.com/gogo/protobuf v1.3.2
|
github.com/gogo/protobuf v1.3.2
|
||||||
github.com/google/gopacket v1.1.19
|
github.com/google/gopacket v1.1.19
|
||||||
github.com/kardianos/service v1.2.4
|
github.com/kardianos/service v1.3.0
|
||||||
github.com/miekg/dns v1.1.72
|
github.com/miekg/dns v1.1.72
|
||||||
github.com/miekg/pkcs11 v1.1.2
|
github.com/miekg/pkcs11 v1.1.2
|
||||||
github.com/nbrownus/go-metrics-prometheus v0.0.0-20210712211119-974a6260965f
|
github.com/nbrownus/go-metrics-prometheus v0.0.0-20210712211119-974a6260965f
|
||||||
@@ -24,15 +24,15 @@ require (
|
|||||||
github.com/vishvananda/netlink v1.3.1
|
github.com/vishvananda/netlink v1.3.1
|
||||||
go.uber.org/goleak v1.3.0
|
go.uber.org/goleak v1.3.0
|
||||||
go.yaml.in/yaml/v3 v3.0.4
|
go.yaml.in/yaml/v3 v3.0.4
|
||||||
golang.org/x/crypto v0.50.0
|
golang.org/x/crypto v0.53.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.56.0
|
||||||
golang.org/x/sync v0.20.0
|
golang.org/x/sync v0.21.0
|
||||||
golang.org/x/sys v0.43.0
|
golang.org/x/sys v0.46.0
|
||||||
golang.org/x/term v0.42.0
|
golang.org/x/term v0.44.0
|
||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
|
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
|
||||||
golang.zx2c4.com/wireguard/windows v0.6.1
|
golang.zx2c4.com/wireguard/windows v1.0.1
|
||||||
google.golang.org/protobuf v1.36.11
|
google.golang.org/protobuf v1.36.11
|
||||||
gopkg.in/yaml.v3 v3.0.1
|
gopkg.in/yaml.v3 v3.0.1
|
||||||
gvisor.dev/gvisor v0.0.0-20240423190808-9d7a357edefe
|
gvisor.dev/gvisor v0.0.0-20240423190808-9d7a357edefe
|
||||||
@@ -50,7 +50,7 @@ require (
|
|||||||
github.com/prometheus/procfs v0.16.1 // indirect
|
github.com/prometheus/procfs v0.16.1 // indirect
|
||||||
github.com/vishvananda/netns v0.0.5 // indirect
|
github.com/vishvananda/netns v0.0.5 // indirect
|
||||||
go.yaml.in/yaml/v2 v2.4.2 // indirect
|
go.yaml.in/yaml/v2 v2.4.2 // indirect
|
||||||
golang.org/x/mod v0.34.0 // indirect
|
golang.org/x/mod v0.36.0 // indirect
|
||||||
golang.org/x/time v0.5.0 // indirect
|
golang.org/x/time v0.5.0 // indirect
|
||||||
golang.org/x/tools v0.43.0 // indirect
|
golang.org/x/tools v0.45.0 // indirect
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -26,8 +26,8 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c
|
|||||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/flynn/noise v1.1.0 h1:KjPQoQCEFdZDiP03phOvGi11+SVVhBG2wOWAorLsstg=
|
github.com/flynn/noise v1.1.0 h1:KjPQoQCEFdZDiP03phOvGi11+SVVhBG2wOWAorLsstg=
|
||||||
github.com/flynn/noise v1.1.0/go.mod h1:xbMo+0i6+IGbYdJhF31t2eR1BIU0CYc12+BNAKwUTag=
|
github.com/flynn/noise v1.1.0/go.mod h1:xbMo+0i6+IGbYdJhF31t2eR1BIU0CYc12+BNAKwUTag=
|
||||||
github.com/gaissmai/bart v0.26.0 h1:xOZ57E9hJLBiQaSyeZa9wgWhGuzfGACgqp4BE77OkO0=
|
github.com/gaissmai/bart v0.28.0 h1:89yZLo8NmyqD0RYgJ3QO9HhqqGGw+oWhf90cZm69Lko=
|
||||||
github.com/gaissmai/bart v0.26.0/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c=
|
github.com/gaissmai/bart v0.28.0/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c=
|
||||||
github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||||
github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||||
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
||||||
@@ -66,8 +66,8 @@ github.com/json-iterator/go v1.1.10/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/
|
|||||||
github.com/json-iterator/go v1.1.11/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4=
|
github.com/json-iterator/go v1.1.11/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4=
|
||||||
github.com/julienschmidt/httprouter v1.2.0/go.mod h1:SYymIcj16QtmaHHD7aYtjjsJG7VTCxuUUipMqKk8s4w=
|
github.com/julienschmidt/httprouter v1.2.0/go.mod h1:SYymIcj16QtmaHHD7aYtjjsJG7VTCxuUUipMqKk8s4w=
|
||||||
github.com/julienschmidt/httprouter v1.3.0/go.mod h1:JR6WtHb+2LUe8TCKY3cZOxFyyO8IZAc4RVcycCCAKdM=
|
github.com/julienschmidt/httprouter v1.3.0/go.mod h1:JR6WtHb+2LUe8TCKY3cZOxFyyO8IZAc4RVcycCCAKdM=
|
||||||
github.com/kardianos/service v1.2.4 h1:XNlGtZOYNx2u91urOdg/Kfmc+gfmuIo1Dd3rEi2OgBk=
|
github.com/kardianos/service v1.3.0 h1:/LGy+xPP2TM+GLTiCZ2di7cy0Jd/qrawlTUfqKYFdTI=
|
||||||
github.com/kardianos/service v1.2.4/go.mod h1:E4V9ufUuY82F7Ztlu1eN9VXWIQxg8NoLQlmFe0MtrXc=
|
github.com/kardianos/service v1.3.0/go.mod h1:E4V9ufUuY82F7Ztlu1eN9VXWIQxg8NoLQlmFe0MtrXc=
|
||||||
github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8=
|
github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8=
|
||||||
github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck=
|
github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck=
|
||||||
github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo=
|
github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo=
|
||||||
@@ -162,16 +162,16 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
|
|||||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||||
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||||
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
||||||
golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
|
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
|
||||||
golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
|
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
|
||||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||||
golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg=
|
golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg=
|
||||||
golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||||
golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||||
golang.org/x/mod v0.34.0 h1:xIHgNUUnW6sYkcM5Jleh05DvLOtwc6RitGHbDk4akRI=
|
golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4=
|
||||||
golang.org/x/mod v0.34.0/go.mod h1:ykgH52iCZe79kzLLMhyCUzhMci+nQj+0XkbXpNYtVjY=
|
golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ=
|
||||||
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||||
golang.org/x/net v0.0.0-20181114220301-adae6a3d119a/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
golang.org/x/net v0.0.0-20181114220301-adae6a3d119a/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||||
golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||||
@@ -182,8 +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.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
||||||
golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw=
|
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
|
||||||
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
||||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
@@ -191,8 +191,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
|||||||
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=
|
||||||
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||||
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
@@ -208,11 +208,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w
|
|||||||
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI=
|
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
|
||||||
golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||||
golang.org/x/term v0.42.0 h1:UiKe+zDFmJobeJ5ggPwOshJIVt6/Ft0rcfrXZDLWAWY=
|
golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
|
||||||
golang.org/x/term v0.42.0/go.mod h1:Dq/D+snpsbazcBG5+F9Q1n2rXV8Ma+71xEjTRufARgY=
|
golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y=
|
||||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||||
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
|
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
|
||||||
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||||
@@ -223,8 +223,8 @@ golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtn
|
|||||||
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
||||||
golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE=
|
golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE=
|
||||||
golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA=
|
golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA=
|
||||||
golang.org/x/tools v0.43.0 h1:12BdW9CeB3Z+J/I/wj34VMl8X+fEXBxVR90JeMX5E7s=
|
golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8=
|
||||||
golang.org/x/tools v0.43.0/go.mod h1:uHkMso649BX2cZK6+RpuIPXS3ho2hZo4FVwfoy1vIk0=
|
golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0=
|
||||||
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
@@ -233,8 +233,8 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu
|
|||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b h1:J1CaxgLerRR5lgx3wnr6L04cJFbWoceSK9JWBdglINo=
|
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b h1:J1CaxgLerRR5lgx3wnr6L04cJFbWoceSK9JWBdglINo=
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b/go.mod h1:tqur9LnfstdR9ep2LaJT4lFUl0EjlHtge+gAjmsHUG4=
|
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b/go.mod h1:tqur9LnfstdR9ep2LaJT4lFUl0EjlHtge+gAjmsHUG4=
|
||||||
golang.zx2c4.com/wireguard/windows v0.6.1 h1:XMaKojH1Hs/raMrmnir4n35nTvzvWj7NmSYzHn2F4qU=
|
golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8=
|
||||||
golang.zx2c4.com/wireguard/windows v0.6.1/go.mod h1:04aqInu5GYuTFvMuDw/rKBAF7mHrltW/3rekpfbbZDM=
|
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
|
||||||
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
|
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
|
||||||
google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=
|
google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=
|
||||||
google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0=
|
google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0=
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ var (
|
|||||||
ErrUnknownSubtype = errors.New("unknown handshake subtype")
|
ErrUnknownSubtype = errors.New("unknown handshake subtype")
|
||||||
ErrMissingContent = errors.New("expected handshake content but message was empty")
|
ErrMissingContent = errors.New("expected handshake content but message was empty")
|
||||||
ErrUnexpectedContent = errors.New("received unexpected handshake content")
|
ErrUnexpectedContent = errors.New("received unexpected handshake content")
|
||||||
|
ErrInvalidRemoteIndex = errors.New("peer sent an invalid index in handshake payload")
|
||||||
ErrIndexAllocation = errors.New("failed to allocate local index")
|
ErrIndexAllocation = errors.New("failed to allocate local index")
|
||||||
ErrNoCredential = errors.New("no handshake credential available for cert version")
|
ErrNoCredential = errors.New("no handshake credential available for cert version")
|
||||||
ErrAsymmetricCipherKeys = errors.New("noise produced only one cipher key")
|
ErrAsymmetricCipherKeys = errors.New("noise produced only one cipher key")
|
||||||
|
|||||||
+12
-2
@@ -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
|
||||||
}
|
}
|
||||||
@@ -310,11 +312,19 @@ func (m *Machine) processPayload(msg []byte, flags msgFlags) error {
|
|||||||
|
|
||||||
// Process payload
|
// Process payload
|
||||||
if flags.expectsPayload {
|
if flags.expectsPayload {
|
||||||
|
var remoteIndex uint32
|
||||||
if m.result.Initiator {
|
if m.result.Initiator {
|
||||||
m.result.RemoteIndex = payload.ResponderIndex
|
remoteIndex = payload.ResponderIndex
|
||||||
} else {
|
} else {
|
||||||
m.result.RemoteIndex = payload.InitiatorIndex
|
remoteIndex = payload.InitiatorIndex
|
||||||
}
|
}
|
||||||
|
// The payload presence check above can be satisfied by Time alone, so a payload
|
||||||
|
// could still carry a zero index here. We need to reject it.
|
||||||
|
if remoteIndex == 0 {
|
||||||
|
m.failed = true
|
||||||
|
return ErrInvalidRemoteIndex
|
||||||
|
}
|
||||||
|
m.result.RemoteIndex = remoteIndex
|
||||||
m.result.HandshakeTime = payload.Time
|
m.result.HandshakeTime = payload.Time
|
||||||
m.payloadSet = true
|
m.payloadSet = true
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -229,6 +229,24 @@ func TestMachineProcessPayload(t *testing.T) {
|
|||||||
require.ErrorIs(t, err, ErrUnexpectedContent)
|
require.ErrorIs(t, err, ErrUnexpectedContent)
|
||||||
assert.True(t, m.Failed())
|
assert.True(t, m.Failed())
|
||||||
})
|
})
|
||||||
|
|
||||||
|
t.Run("zero initiator index on responder is fatal", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
bytes := MarshalPayload(nil, Payload{InitiatorIndex: 0, Time: 1})
|
||||||
|
err := m.processPayload(bytes, msgFlags{expectsPayload: true})
|
||||||
|
require.ErrorIs(t, err, ErrInvalidRemoteIndex)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
assert.Zero(t, m.result.RemoteIndex)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("zero responder index on initiator is fatal", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, true, 100)
|
||||||
|
bytes := MarshalPayload(nil, Payload{InitiatorIndex: 100, ResponderIndex: 0, Time: 1})
|
||||||
|
err := m.processPayload(bytes, msgFlags{expectsPayload: true})
|
||||||
|
require.ErrorIs(t, err, ErrInvalidRemoteIndex)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
assert.Zero(t, m.result.RemoteIndex)
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestMachineRequireComplete checks the fail-on-incomplete-handshake path
|
// TestMachineRequireComplete checks the fail-on-incomplete-handshake path
|
||||||
|
|||||||
+5
-11
@@ -83,6 +83,7 @@ type HandshakeHostInfo struct {
|
|||||||
initiatingVersionOverride cert.Version // Should we use a non-default cert version for this handshake?
|
initiatingVersionOverride cert.Version // Should we use a non-default cert version for this handshake?
|
||||||
counter int64 // How many attempts have we made so far
|
counter int64 // How many attempts have we made so far
|
||||||
lastRemotes []netip.AddrPort // Remotes that we sent to during the previous attempt
|
lastRemotes []netip.AddrPort // Remotes that we sent to during the previous attempt
|
||||||
|
lastRelays []netip.Addr // Relays we attempted to use during the previous attempt
|
||||||
packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes
|
packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes
|
||||||
|
|
||||||
hostinfo *HostInfo
|
hostinfo *HostInfo
|
||||||
@@ -217,7 +218,6 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
fields := []any{
|
fields := []any{
|
||||||
"udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()),
|
"udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()),
|
||||||
"initiatorIndex", hh.hostinfo.localIndexId,
|
"initiatorIndex", hh.hostinfo.localIndexId,
|
||||||
"remoteIndex", hh.hostinfo.remoteIndexId,
|
|
||||||
"durationNs", time.Since(hh.startTime).Nanoseconds(),
|
"durationNs", time.Since(hh.startTime).Nanoseconds(),
|
||||||
}
|
}
|
||||||
// hh.machine can be nil here if buildStage0Packet never succeeded
|
// hh.machine can be nil here if buildStage0Packet never succeeded
|
||||||
@@ -323,7 +323,7 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
hm.f.relayManager.StartRelays(hm.f, vpnIp, hostinfo, stage0)
|
hm.f.relayManager.StartRelays(hm.f, vpnIp, hh, stage0)
|
||||||
|
|
||||||
// If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add
|
// If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add
|
||||||
if !lighthouseTriggered {
|
if !lighthouseTriggered {
|
||||||
@@ -430,14 +430,11 @@ func (hm *HandshakeManager) CheckAndComplete(hostinfo *HostInfo, handshakePacket
|
|||||||
// Check if we already have a tunnel with this vpn ip
|
// Check if we already have a tunnel with this vpn ip
|
||||||
existingHostInfo, found := hm.mainHostMap.Hosts[hostinfo.vpnAddrs[0]]
|
existingHostInfo, found := hm.mainHostMap.Hosts[hostinfo.vpnAddrs[0]]
|
||||||
if found && existingHostInfo != nil {
|
if found && existingHostInfo != nil {
|
||||||
testHostInfo := existingHostInfo
|
// Is it just a delayed handshake packet? Check every hostinfo we hold for this address.
|
||||||
for testHostInfo != nil {
|
for _, testHostInfo := range hm.mainHostMap.unlockedGetHostList(hostinfo.vpnAddrs[0]) {
|
||||||
// Is it just a delayed handshake packet?
|
|
||||||
if bytes.Equal(hostinfo.HandshakePacket[handshakePacket], testHostInfo.HandshakePacket[handshakePacket]) {
|
if bytes.Equal(hostinfo.HandshakePacket[handshakePacket], testHostInfo.HandshakePacket[handshakePacket]) {
|
||||||
return testHostInfo, ErrAlreadySeen
|
return testHostInfo, ErrAlreadySeen
|
||||||
}
|
}
|
||||||
|
|
||||||
testHostInfo = testHostInfo.next
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Is this a newer handshake?
|
// Is this a newer handshake?
|
||||||
@@ -465,7 +462,6 @@ func (hm *HandshakeManager) CheckAndComplete(hostinfo *HostInfo, handshakePacket
|
|||||||
// We have a collision, but this can happen since we can't control
|
// We have a collision, but this can happen since we can't control
|
||||||
// the remote ID. Just log about the situation as a note.
|
// the remote ID. Just log about the situation as a note.
|
||||||
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
||||||
"remoteIndex", hostinfo.remoteIndexId,
|
|
||||||
"collision", existingRemoteIndex.vpnAddrs,
|
"collision", existingRemoteIndex.vpnAddrs,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -488,7 +484,6 @@ func (hm *HandshakeManager) Complete(hostinfo *HostInfo, f *Interface) {
|
|||||||
// We have a collision, but this can happen since we can't control
|
// We have a collision, but this can happen since we can't control
|
||||||
// the remote ID. Just log about the situation as a note.
|
// the remote ID. Just log about the situation as a note.
|
||||||
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
||||||
"remoteIndex", hostinfo.remoteIndexId,
|
|
||||||
"collision", existingRemoteIndex.vpnAddrs,
|
"collision", existingRemoteIndex.vpnAddrs,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -798,7 +793,6 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head
|
|||||||
}
|
}
|
||||||
|
|
||||||
hm.sendHandshakeResponse(via, response, hostinfo, false)
|
hm.sendHandshakeResponse(via, response, hostinfo, false)
|
||||||
f.connectionManager.AddTrafficWatch(hostinfo)
|
|
||||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||||
|
|
||||||
// Don't wait for UpdateWorker
|
// Don't wait for UpdateWorker
|
||||||
@@ -965,7 +959,6 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
|||||||
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
||||||
|
|
||||||
hm.Complete(hostinfo, f)
|
hm.Complete(hostinfo, f)
|
||||||
f.connectionManager.AddTrafficWatch(hostinfo)
|
|
||||||
|
|
||||||
if len(hh.packetStore) > 0 {
|
if len(hh.packetStore) > 0 {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
@@ -974,6 +967,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)
|
||||||
|
|||||||
+171
-118
@@ -56,11 +56,20 @@ type Relay struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type HostMap struct {
|
type HostMap struct {
|
||||||
sync.RWMutex //Because we concurrently read and write to our maps
|
sync.RWMutex //Because we concurrently read and write to our maps
|
||||||
Indexes map[uint32]*HostInfo
|
Indexes map[uint32]*HostInfo
|
||||||
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
|
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
|
||||||
RemoteIndexes map[uint32]*HostInfo
|
RemoteIndexes map[uint32]*HostInfo
|
||||||
|
// Hosts maps a vpn address to its primary hostinfo, one entry per address we hold a tunnel
|
||||||
|
// for. moreHosts only has an entry while an address is held by 2 or more hostinfos and stores
|
||||||
|
// the full most-recent-first list; moreHosts[a][0] is always the same hostinfo as Hosts[a].
|
||||||
|
// Each address gets its own independent list, so a hostinfo owning multiple addresses can
|
||||||
|
// never corrupt another address's ordering the way the old shared next/prev chain could.
|
||||||
|
// Entries in moreHosts are only ever written by unlockedSetHostsForAddr; Hosts is written
|
||||||
|
// directly only in the single-hostinfo fast paths where moreHosts is known to have no entry,
|
||||||
|
// and unlockedDeleteHostInfo swaps either map for a fresh one when it fully drains.
|
||||||
Hosts map[netip.Addr]*HostInfo
|
Hosts map[netip.Addr]*HostInfo
|
||||||
|
moreHosts map[netip.Addr][]*HostInfo
|
||||||
preferredRanges atomic.Pointer[[]netip.Prefix]
|
preferredRanges atomic.Pointer[[]netip.Prefix]
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
@@ -138,9 +147,9 @@ func (rs *RelayState) InsertRelayTo(ip netip.Addr) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (rs *RelayState) CopyRelayIps() []netip.Addr {
|
func (rs *RelayState) CopyRelayIps() []netip.Addr {
|
||||||
ret := make([]netip.Addr, len(rs.relays))
|
|
||||||
rs.RLock()
|
rs.RLock()
|
||||||
defer rs.RUnlock()
|
defer rs.RUnlock()
|
||||||
|
ret := make([]netip.Addr, len(rs.relays))
|
||||||
copy(ret, rs.relays)
|
copy(ret, rs.relays)
|
||||||
return ret
|
return ret
|
||||||
}
|
}
|
||||||
@@ -229,7 +238,7 @@ const (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type HostInfo struct {
|
type HostInfo struct {
|
||||||
remote netip.AddrPort
|
remote atomic.Pointer[netip.AddrPort]
|
||||||
remotes *RemoteList
|
remotes *RemoteList
|
||||||
promoteCounter atomic.Uint32
|
promoteCounter atomic.Uint32
|
||||||
ConnectionState *ConnectionState
|
ConnectionState *ConnectionState
|
||||||
@@ -266,10 +275,6 @@ type HostInfo struct {
|
|||||||
lastRoam time.Time
|
lastRoam time.Time
|
||||||
lastRoamRemote netip.AddrPort
|
lastRoamRemote netip.AddrPort
|
||||||
|
|
||||||
// Used to track other hostinfos for this vpn ip since only 1 can be primary
|
|
||||||
// Synchronised via hostmap lock and not the hostinfo lock.
|
|
||||||
next, prev *HostInfo
|
|
||||||
|
|
||||||
//TODO: in, out, and others might benefit from being an atomic.Int32. We could collapse connectionManager pendingDeletion, relayUsed, and in/out into this 1 thing
|
//TODO: in, out, and others might benefit from being an atomic.Int32. We could collapse connectionManager pendingDeletion, relayUsed, and in/out into this 1 thing
|
||||||
in, out, pendingDeletion atomic.Bool
|
in, out, pendingDeletion atomic.Bool
|
||||||
|
|
||||||
@@ -334,6 +339,7 @@ func newHostMap(l *slog.Logger) *HostMap {
|
|||||||
Relays: map[uint32]*HostInfo{},
|
Relays: map[uint32]*HostInfo{},
|
||||||
RemoteIndexes: map[uint32]*HostInfo{},
|
RemoteIndexes: map[uint32]*HostInfo{},
|
||||||
Hosts: map[netip.Addr]*HostInfo{},
|
Hosts: map[netip.Addr]*HostInfo{},
|
||||||
|
moreHosts: map[netip.Addr][]*HostInfo{},
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -382,13 +388,55 @@ func (hm *HostMap) EmitStats() {
|
|||||||
metrics.GetOrRegisterGauge("hostmap.main.relayIndexes", nil).Update(int64(relaysLen))
|
metrics.GetOrRegisterGauge("hostmap.main.relayIndexes", nil).Update(int64(relaysLen))
|
||||||
}
|
}
|
||||||
|
|
||||||
// DeleteHostInfo will fully unlink the hostinfo and return true if it was the final hostinfo for this vpn ip
|
// unlockedSetHostsForAddr stores the per-address hostinfo list (list[0] is the primary). An empty
|
||||||
|
// list removes the address. This is the one place Hosts and moreHosts are written together, keep
|
||||||
|
// it that way. Callers must hold the write lock.
|
||||||
|
func (hm *HostMap) unlockedSetHostsForAddr(addr netip.Addr, list []*HostInfo) {
|
||||||
|
if len(list) == 0 {
|
||||||
|
delete(hm.Hosts, addr)
|
||||||
|
delete(hm.moreHosts, addr)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
hm.Hosts[addr] = list[0]
|
||||||
|
if len(list) > 1 {
|
||||||
|
hm.moreHosts[addr] = list
|
||||||
|
} else {
|
||||||
|
delete(hm.moreHosts, addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// unlockedGetHostList returns every hostinfo holding addr, primary first, or nil if we have no
|
||||||
|
// tunnel for addr. The common single-hostinfo case builds a fresh one element list, so keep this
|
||||||
|
// off the packet hot path; the primary is a direct Hosts read. Callers must hold the lock (read
|
||||||
|
// or write).
|
||||||
|
func (hm *HostMap) unlockedGetHostList(addr netip.Addr) []*HostInfo {
|
||||||
|
if list, ok := hm.moreHosts[addr]; ok {
|
||||||
|
return list
|
||||||
|
}
|
||||||
|
if h, ok := hm.Hosts[addr]; ok {
|
||||||
|
return []*HostInfo{h}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// removeHostInfo returns list with hi removed (order preserved), or list unchanged if hi is
|
||||||
|
// absent. It deletes in place: every mutator holds the hostmap write lock and no reader ever
|
||||||
|
// retains a slice across a mutation (readers iterate under RLock), so there is no snapshot to
|
||||||
|
// invalidate.
|
||||||
|
func removeHostInfo(list []*HostInfo, hi *HostInfo) []*HostInfo {
|
||||||
|
idx := slices.Index(list, hi)
|
||||||
|
if idx < 0 {
|
||||||
|
return list
|
||||||
|
}
|
||||||
|
return slices.Delete(list, idx, idx+1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteHostInfo will fully unlink the hostinfo and return true if no other hostinfo still holds
|
||||||
|
// any of its vpn addrs, meaning we no longer have a tunnel to the peer
|
||||||
func (hm *HostMap) DeleteHostInfo(hostinfo *HostInfo) bool {
|
func (hm *HostMap) DeleteHostInfo(hostinfo *HostInfo) bool {
|
||||||
// Delete the host itself, ensuring it's not modified anymore
|
// Delete the host itself, ensuring it's not modified anymore
|
||||||
hm.Lock()
|
hm.Lock()
|
||||||
// If we have a previous or next hostinfo then we are not the last one for this vpn ip
|
final := hm.unlockedDeleteHostInfo(hostinfo)
|
||||||
final := (hostinfo.next == nil && hostinfo.prev == nil)
|
|
||||||
hm.unlockedDeleteHostInfo(hostinfo)
|
|
||||||
hm.Unlock()
|
hm.Unlock()
|
||||||
|
|
||||||
return final
|
return final
|
||||||
@@ -400,85 +448,66 @@ func (hm *HostMap) MakePrimary(hostinfo *HostInfo) {
|
|||||||
hm.unlockedMakePrimary(hostinfo)
|
hm.unlockedMakePrimary(hostinfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) {
|
// unlockedMakePrimary reports whether hostinfo is (now) the primary for each of its addresses,
|
||||||
// Get the current primary, if it exists
|
// false only when it is no longer in the hostmap at all.
|
||||||
oldHostinfo := hm.Hosts[hostinfo.vpnAddrs[0]]
|
func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) bool {
|
||||||
|
// A hostinfo that is no longer in the hostmap must not be re-inserted here. Callers can race
|
||||||
// Every address in the hostinfo gets elevated to primary
|
// tunnel teardown, deciding to promote under the read lock and only taking the write lock
|
||||||
for _, vpnAddr := range hostinfo.vpnAddrs {
|
// after a delete fully unlinked the hostinfo (connection manager swapPrimary, AddRelay). Every
|
||||||
//NOTE: It is possible that we leave a dangling hostinfo here but connection manager works on
|
// live hostinfo is registered in Indexes by unlockedAddHostInfo, so this is a membership test.
|
||||||
// indexes so it should be fine.
|
if hm.Indexes[hostinfo.localIndexId] != hostinfo {
|
||||||
hm.Hosts[vpnAddr] = hostinfo
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// If we are already primary then we won't bother re-linking
|
// Move hostinfo to the front (primary) of each of its address lists. The lists are
|
||||||
if oldHostinfo == hostinfo {
|
// independent per address, so this can never leave a dangling entry the way promoting
|
||||||
return
|
// against a single shared chain could.
|
||||||
}
|
|
||||||
|
|
||||||
// Unlink this hostinfo
|
|
||||||
if hostinfo.prev != nil {
|
|
||||||
hostinfo.prev.next = hostinfo.next
|
|
||||||
}
|
|
||||||
if hostinfo.next != nil {
|
|
||||||
hostinfo.next.prev = hostinfo.prev
|
|
||||||
}
|
|
||||||
|
|
||||||
// If there wasn't a previous primary then clear out any links
|
|
||||||
if oldHostinfo == nil {
|
|
||||||
hostinfo.next = nil
|
|
||||||
hostinfo.prev = nil
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Relink the hostinfo as primary
|
|
||||||
hostinfo.next = oldHostinfo
|
|
||||||
oldHostinfo.prev = hostinfo
|
|
||||||
hostinfo.prev = nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
|
||||||
for _, addr := range hostinfo.vpnAddrs {
|
for _, addr := range hostinfo.vpnAddrs {
|
||||||
h := hm.Hosts[addr]
|
if hm.Hosts[addr] == hostinfo {
|
||||||
for h != nil {
|
// Already primary for this address, the list is already in the right order
|
||||||
if h == hostinfo {
|
continue
|
||||||
hm.unlockedInnerDeleteHostInfo(h, addr)
|
|
||||||
}
|
|
||||||
h = h.next
|
|
||||||
}
|
}
|
||||||
|
list := removeHostInfo(hm.unlockedGetHostList(addr), hostinfo)
|
||||||
|
list = append([]*HostInfo{hostinfo}, list...)
|
||||||
|
hm.unlockedSetHostsForAddr(addr, list)
|
||||||
}
|
}
|
||||||
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HostMap) unlockedInnerDeleteHostInfo(hostinfo *HostInfo, addr netip.Addr) {
|
// unlockedDeleteHostInfo removes hostinfo from every one of its address lists and from the index
|
||||||
primary, ok := hm.Hosts[addr]
|
// maps. It returns true if this was the last hostinfo for all of its addresses (we no longer have
|
||||||
isLastHostinfo := hostinfo.next == nil && hostinfo.prev == nil
|
// any tunnel to the peer), which the caller uses to decide whether to clear learned lighthouse
|
||||||
if ok && primary == hostinfo {
|
// state and disestablish relays.
|
||||||
// The vpn addr pointer points to the same hostinfo as the local index id, we can remove it
|
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
||||||
delete(hm.Hosts, addr)
|
// Remove this hostinfo from each of its address lists. The lists are independent, so a
|
||||||
if len(hm.Hosts) == 0 {
|
// sibling is never promoted to an address it does not own and no other list is touched.
|
||||||
hm.Hosts = map[netip.Addr]*HostInfo{}
|
final := true
|
||||||
}
|
for _, addr := range hostinfo.vpnAddrs {
|
||||||
|
if list, ok := hm.moreHosts[addr]; ok {
|
||||||
if hostinfo.next != nil {
|
list = removeHostInfo(list, hostinfo)
|
||||||
// We had more than 1 hostinfo at this vpn addr, promote the next in the list to primary
|
hm.unlockedSetHostsForAddr(addr, list)
|
||||||
hm.Hosts[addr] = hostinfo.next
|
if len(list) > 0 {
|
||||||
// It is primary, there is no previous hostinfo now
|
final = false
|
||||||
hostinfo.next.prev = nil
|
}
|
||||||
}
|
} else if existing, ok := hm.Hosts[addr]; ok {
|
||||||
|
if existing == hostinfo {
|
||||||
} else {
|
// Common case, the only hostinfo for this address. moreHosts has no entry to clean up.
|
||||||
// Relink if we were in the middle of multiple hostinfos for this vpn addr
|
delete(hm.Hosts, addr)
|
||||||
if hostinfo.prev != nil {
|
} else {
|
||||||
hostinfo.prev.next = hostinfo.next
|
// We don't hold this address but another hostinfo does, we still have a tunnel to the peer
|
||||||
}
|
final = false
|
||||||
|
}
|
||||||
if hostinfo.next != nil {
|
|
||||||
hostinfo.next.prev = hostinfo.prev
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
hostinfo.next = nil
|
// Go maps never shrink their buckets, replace fully drained maps so a node that churned
|
||||||
hostinfo.prev = nil
|
// through a large peer count gives the memory back. Same idiom as the index maps below.
|
||||||
|
if len(hm.Hosts) == 0 {
|
||||||
|
hm.Hosts = map[netip.Addr]*HostInfo{}
|
||||||
|
}
|
||||||
|
if len(hm.moreHosts) == 0 {
|
||||||
|
hm.moreHosts = map[netip.Addr][]*HostInfo{}
|
||||||
|
}
|
||||||
|
|
||||||
// The remote index uses index ids outside our control so lets make sure we are only removing
|
// The remote index uses index ids outside our control so lets make sure we are only removing
|
||||||
// the remote index pointer here if it points to the hostinfo we are deleting
|
// the remote index pointer here if it points to the hostinfo we are deleting
|
||||||
@@ -502,7 +531,7 @@ func (hm *HostMap) unlockedInnerDeleteHostInfo(hostinfo *HostInfo, addr netip.Ad
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
if isLastHostinfo {
|
if final {
|
||||||
// I have lost connectivity to my peers. My relay tunnel is likely broken. Mark the next
|
// I have lost connectivity to my peers. My relay tunnel is likely broken. Mark the next
|
||||||
// hops as 'Requested' so that new relay tunnels are created in the future.
|
// hops as 'Requested' so that new relay tunnels are created in the future.
|
||||||
hm.unlockedDisestablishVpnAddrRelayFor(hostinfo)
|
hm.unlockedDisestablishVpnAddrRelayFor(hostinfo)
|
||||||
@@ -511,6 +540,8 @@ func (hm *HostMap) unlockedInnerDeleteHostInfo(hostinfo *HostInfo, addr netip.Ad
|
|||||||
for _, localRelayIdx := range hostinfo.relayState.CopyRelayForIdxs() {
|
for _, localRelayIdx := range hostinfo.relayState.CopyRelayForIdxs() {
|
||||||
delete(hm.Relays, localRelayIdx)
|
delete(hm.Relays, localRelayIdx)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return final
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
|
func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
|
||||||
@@ -554,19 +585,30 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net
|
|||||||
hm.RLock()
|
hm.RLock()
|
||||||
defer hm.RUnlock()
|
defer hm.RUnlock()
|
||||||
|
|
||||||
|
// This runs per relayed packet, so check the primary with a single map probe and only consult
|
||||||
|
// moreHosts when the primary can't relay for us.
|
||||||
h, ok := hm.Hosts[relayHostIp]
|
h, ok := hm.Hosts[relayHostIp]
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, nil, errors.New("unable to find host")
|
return nil, nil, errors.New("unable to find host")
|
||||||
}
|
}
|
||||||
|
|
||||||
for h != nil {
|
for _, targetIp := range targetIps {
|
||||||
for _, targetIp := range targetIps {
|
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
||||||
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
if ok && r.State == Established {
|
||||||
if ok && r.State == Established {
|
return h, r, nil
|
||||||
return h, r, nil
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if list, ok := hm.moreHosts[relayHostIp]; ok {
|
||||||
|
// list[0] is the primary we already checked
|
||||||
|
for _, h := range list[1:] {
|
||||||
|
for _, targetIp := range targetIps {
|
||||||
|
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
||||||
|
if ok && r.State == Established {
|
||||||
|
return h, r, nil
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
h = h.next
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil, nil, errors.New("unable to find host with relay")
|
return nil, nil, errors.New("unable to find host with relay")
|
||||||
@@ -574,20 +616,14 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net
|
|||||||
|
|
||||||
func (hm *HostMap) unlockedDisestablishVpnAddrRelayFor(hi *HostInfo) {
|
func (hm *HostMap) unlockedDisestablishVpnAddrRelayFor(hi *HostInfo) {
|
||||||
for _, relayHostIp := range hi.relayState.CopyRelayIps() {
|
for _, relayHostIp := range hi.relayState.CopyRelayIps() {
|
||||||
if h, ok := hm.Hosts[relayHostIp]; ok {
|
for _, h := range hm.unlockedGetHostList(relayHostIp) {
|
||||||
for h != nil {
|
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
||||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
|
||||||
h = h.next
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for _, rs := range hi.relayState.CopyAllRelayFor() {
|
for _, rs := range hi.relayState.CopyAllRelayFor() {
|
||||||
if rs.Type == ForwardingType {
|
if rs.Type == ForwardingType {
|
||||||
if h, ok := hm.Hosts[rs.PeerAddr]; ok {
|
for _, h := range hm.unlockedGetHostList(rs.PeerAddr) {
|
||||||
for h != nil {
|
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
||||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
|
||||||
h = h.next
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -623,6 +659,11 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
|
|||||||
hm.Indexes[hostinfo.localIndexId] = hostinfo
|
hm.Indexes[hostinfo.localIndexId] = hostinfo
|
||||||
hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo
|
hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo
|
||||||
|
|
||||||
|
hostinfo.out.Store(true)
|
||||||
|
if f.connectionManager != nil { // f.connectionManager is only nil in some unit tests
|
||||||
|
f.connectionManager.trafficTimer.Add(hostinfo.localIndexId, f.connectionManager.checkInterval)
|
||||||
|
}
|
||||||
|
|
||||||
if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hm.l.Debug("Hostmap vpnIp added",
|
hm.l.Debug("Hostmap vpnIp added",
|
||||||
"hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts),
|
"hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts),
|
||||||
@@ -632,22 +673,27 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) {
|
func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) {
|
||||||
existing := hm.Hosts[vpnAddr]
|
existing, ok := hm.Hosts[vpnAddr]
|
||||||
hm.Hosts[vpnAddr] = hostinfo
|
if !ok {
|
||||||
|
// Common case, the first hostinfo for this address. moreHosts stays empty.
|
||||||
if existing != nil && existing != hostinfo {
|
hm.Hosts[vpnAddr] = hostinfo
|
||||||
hostinfo.next = existing
|
return
|
||||||
existing.prev = hostinfo
|
|
||||||
}
|
}
|
||||||
|
|
||||||
i := 1
|
// The new hostinfo becomes the primary for this address. Remove any stale copy of it first so
|
||||||
check := hostinfo
|
// we never hold a duplicate, then prepend.
|
||||||
for check != nil {
|
list, ok := hm.moreHosts[vpnAddr]
|
||||||
if i > MaxHostInfosPerVpnIp {
|
if !ok {
|
||||||
hm.unlockedDeleteHostInfo(check)
|
list = []*HostInfo{existing}
|
||||||
}
|
}
|
||||||
check = check.next
|
list = removeHostInfo(list, hostinfo)
|
||||||
i++
|
list = append([]*HostInfo{hostinfo}, list...)
|
||||||
|
hm.unlockedSetHostsForAddr(vpnAddr, list)
|
||||||
|
|
||||||
|
// Enforce the per-address cap by fully retiring the oldest hostinfo once we exceed it.
|
||||||
|
// Deleting it removes it from all of its addresses and the index maps, matching prior behavior.
|
||||||
|
if len(list) > MaxHostInfosPerVpnIp {
|
||||||
|
hm.unlockedDeleteHostInfo(list[len(list)-1])
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -679,7 +725,7 @@ func (hm *HostMap) ForEachIndex(f controlEach) {
|
|||||||
func (i *HostInfo) TryPromoteBest(preferredRanges []netip.Prefix, ifce *Interface) {
|
func (i *HostInfo) TryPromoteBest(preferredRanges []netip.Prefix, ifce *Interface) {
|
||||||
c := i.promoteCounter.Add(1)
|
c := i.promoteCounter.Add(1)
|
||||||
if c%ifce.tryPromoteEvery.Load() == 0 {
|
if c%ifce.tryPromoteEvery.Load() == 0 {
|
||||||
remote := i.remote
|
remote := i.GetRemote()
|
||||||
|
|
||||||
// return early if we are already on a preferred remote
|
// return early if we are already on a preferred remote
|
||||||
if remote.IsValid() {
|
if remote.IsValid() {
|
||||||
@@ -721,11 +767,18 @@ func (i *HostInfo) GetCert() *cert.CachedCertificate {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (i *HostInfo) GetRemote() netip.AddrPort {
|
||||||
|
if p := i.remote.Load(); p != nil {
|
||||||
|
return *p
|
||||||
|
}
|
||||||
|
return netip.AddrPort{}
|
||||||
|
}
|
||||||
|
|
||||||
// TODO: Maybe use ViaSender here?
|
// TODO: Maybe use ViaSender here?
|
||||||
func (i *HostInfo) SetRemote(remote netip.AddrPort) {
|
func (i *HostInfo) SetRemote(remote netip.AddrPort) {
|
||||||
// We copy here because we likely got this remote from a source that reuses the object
|
// We copy here because we likely got this remote from a source that reuses the object
|
||||||
if i.remote != remote {
|
if i.GetRemote() != remote {
|
||||||
i.remote = remote
|
i.remote.Store(&remote)
|
||||||
i.remotes.LearnRemote(i.vpnAddrs[0], remote)
|
i.remotes.LearnRemote(i.vpnAddrs[0], remote)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -737,7 +790,7 @@ func (i *HostInfo) SetRemoteIfPreferred(hm *HostMap, via ViaSender) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
currentRemote := i.remote
|
currentRemote := i.GetRemote()
|
||||||
if !currentRemote.IsValid() {
|
if !currentRemote.IsValid() {
|
||||||
i.SetRemote(via.UdpAddr)
|
i.SetRemote(via.UdpAddr)
|
||||||
return true
|
return true
|
||||||
|
|||||||
+294
-137
@@ -2,6 +2,7 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"slices"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
@@ -10,78 +11,84 @@ import (
|
|||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// chainIds returns the localIndexIds of the hostinfos holding addr, primary (index 0) first. It
|
||||||
|
// also validates the Hosts/moreHosts sync contract on every call so a mutation that broke it
|
||||||
|
// fails fast.
|
||||||
|
func chainIds(t *testing.T, hm *HostMap, addr netip.Addr) []uint32 {
|
||||||
|
t.Helper()
|
||||||
|
assertHostMapInvariants(t, hm)
|
||||||
|
list := hm.unlockedGetHostList(addr)
|
||||||
|
ids := make([]uint32, len(list))
|
||||||
|
for i, h := range list {
|
||||||
|
ids[i] = h.localIndexId
|
||||||
|
}
|
||||||
|
return ids
|
||||||
|
}
|
||||||
|
|
||||||
|
// assertHostMapInvariants checks the Hosts/moreHosts contract: moreHosts only holds addresses
|
||||||
|
// with 2 or more hostinfos, its first entry is always the primary in Hosts, lists never hold
|
||||||
|
// duplicates, every hostinfo in a list owns the address and is registered in Indexes, and every
|
||||||
|
// indexed hostinfo is reachable through each of its addresses.
|
||||||
|
func assertHostMapInvariants(t *testing.T, hm *HostMap) {
|
||||||
|
t.Helper()
|
||||||
|
for addr, list := range hm.moreHosts {
|
||||||
|
require.GreaterOrEqualf(t, len(list), 2, "moreHosts[%s] must hold at least 2 hostinfos", addr)
|
||||||
|
require.Samef(t, hm.Hosts[addr], list[0], "moreHosts[%s][0] must match the primary in Hosts", addr)
|
||||||
|
seen := map[*HostInfo]bool{}
|
||||||
|
for _, h := range list {
|
||||||
|
require.NotNilf(t, h, "moreHosts[%s] must never hold a nil hostinfo", addr)
|
||||||
|
require.Falsef(t, seen[h], "moreHosts[%s] holds hostinfo %d twice", addr, h.localIndexId)
|
||||||
|
seen[h] = true
|
||||||
|
require.Samef(t, hm.Indexes[h.localIndexId], h, "moreHosts[%s] member %d is not registered in Indexes", addr, h.localIndexId)
|
||||||
|
require.Truef(t, slices.Contains(h.vpnAddrs, addr), "moreHosts[%s] member %d does not own the address", addr, h.localIndexId)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for addr, h := range hm.Hosts {
|
||||||
|
require.NotNilf(t, h, "Hosts[%s] must never be nil", addr)
|
||||||
|
require.Samef(t, hm.Indexes[h.localIndexId], h, "Hosts[%s] primary %d is not registered in Indexes", addr, h.localIndexId)
|
||||||
|
require.Truef(t, slices.Contains(h.vpnAddrs, addr), "Hosts[%s] primary (index %d) does not own the address", addr, h.localIndexId)
|
||||||
|
}
|
||||||
|
for idx, h := range hm.Indexes {
|
||||||
|
require.Equalf(t, idx, h.localIndexId, "Indexes[%d] holds hostinfo with localIndexId %d", idx, h.localIndexId)
|
||||||
|
for _, va := range h.vpnAddrs {
|
||||||
|
require.Truef(t, slices.Contains(hm.unlockedGetHostList(va), h), "indexed hostinfo %d is missing from the list for %s", idx, va)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestHostMap_MakePrimary(t *testing.T) {
|
func TestHostMap_MakePrimary(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
hm := newHostMap(l)
|
hm := newHostMap(l)
|
||||||
|
|
||||||
f := &Interface{}
|
f := &Interface{}
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
|
||||||
h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
|
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
|
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
||||||
h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
|
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3}
|
||||||
h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4}
|
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4}
|
||||||
|
|
||||||
hm.unlockedAddHostInfo(h4, f)
|
hm.unlockedAddHostInfo(h4, f)
|
||||||
hm.unlockedAddHostInfo(h3, f)
|
hm.unlockedAddHostInfo(h3, f)
|
||||||
hm.unlockedAddHostInfo(h2, f)
|
hm.unlockedAddHostInfo(h2, f)
|
||||||
hm.unlockedAddHostInfo(h1, f)
|
hm.unlockedAddHostInfo(h1, f)
|
||||||
|
|
||||||
// Make sure we go h1 -> h2 -> h3 -> h4
|
// Most-recently-added is primary: h1, h2, h3, h4
|
||||||
prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
assert.Equal(t, []uint32{1, 2, 3, 4}, chainIds(t, hm, a))
|
||||||
assert.Equal(t, h1.localIndexId, prim.localIndexId)
|
assert.Equal(t, h1, hm.QueryVpnAddr(a))
|
||||||
assert.Equal(t, h2.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Nil(t, h4.next)
|
|
||||||
|
|
||||||
// Swap h3/middle to primary
|
// Swap the middle to primary: h3, h1, h2, h4
|
||||||
hm.MakePrimary(h3)
|
hm.MakePrimary(h3)
|
||||||
|
assert.Equal(t, []uint32{3, 1, 2, 4}, chainIds(t, hm, a))
|
||||||
|
assert.Equal(t, h3, hm.QueryVpnAddr(a))
|
||||||
|
|
||||||
// Make sure we go h3 -> h1 -> h2 -> h4
|
// Swap the tail to primary: h4, h3, h1, h2
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
|
||||||
assert.Equal(t, h3.localIndexId, prim.localIndexId)
|
|
||||||
assert.Equal(t, h1.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Nil(t, h4.next)
|
|
||||||
|
|
||||||
// Swap h4/tail to primary
|
|
||||||
hm.MakePrimary(h4)
|
hm.MakePrimary(h4)
|
||||||
|
assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a))
|
||||||
|
|
||||||
// Make sure we go h4 -> h3 -> h1 -> h2
|
// Swapping the current primary again is a no-op
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
|
||||||
assert.Equal(t, h4.localIndexId, prim.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h1.localIndexId, h3.next.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h3.prev.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
|
|
||||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
|
||||||
assert.Nil(t, h2.next)
|
|
||||||
|
|
||||||
// Swap h4 again should be no-op
|
|
||||||
hm.MakePrimary(h4)
|
hm.MakePrimary(h4)
|
||||||
|
assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a))
|
||||||
// Make sure we go h4 -> h3 -> h1 -> h2
|
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
|
||||||
assert.Equal(t, h4.localIndexId, prim.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h1.localIndexId, h3.next.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h3.prev.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
|
|
||||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
|
||||||
assert.Nil(t, h2.next)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHostMap_DeleteHostInfo(t *testing.T) {
|
func TestHostMap_DeleteHostInfo(t *testing.T) {
|
||||||
@@ -89,13 +96,14 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
|
|||||||
hm := newHostMap(l)
|
hm := newHostMap(l)
|
||||||
|
|
||||||
f := &Interface{}
|
f := &Interface{}
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
|
||||||
h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
|
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
|
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
||||||
h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
|
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3}
|
||||||
h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4}
|
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4}
|
||||||
h5 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 5}
|
h5 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 5}
|
||||||
h6 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 6}
|
h6 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 6}
|
||||||
|
|
||||||
hm.unlockedAddHostInfo(h6, f)
|
hm.unlockedAddHostInfo(h6, f)
|
||||||
hm.unlockedAddHostInfo(h5, f)
|
hm.unlockedAddHostInfo(h5, f)
|
||||||
@@ -104,94 +112,243 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
|
|||||||
hm.unlockedAddHostInfo(h2, f)
|
hm.unlockedAddHostInfo(h2, f)
|
||||||
hm.unlockedAddHostInfo(h1, f)
|
hm.unlockedAddHostInfo(h1, f)
|
||||||
|
|
||||||
// h6 should be deleted
|
// h6 is evicted by the MaxHostInfosPerVpnIp cap; the rest are newest-first.
|
||||||
assert.Nil(t, h6.next)
|
assert.Nil(t, hm.QueryIndex(h6.localIndexId))
|
||||||
assert.Nil(t, h6.prev)
|
assert.Equal(t, []uint32{1, 2, 3, 4, 5}, chainIds(t, hm, a))
|
||||||
h := hm.QueryIndex(h6.localIndexId)
|
|
||||||
assert.Nil(t, h)
|
|
||||||
|
|
||||||
// Make sure we go h1 -> h2 -> h3 -> h4 -> h5
|
// Delete primary; not final since siblings remain.
|
||||||
prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
assert.False(t, hm.DeleteHostInfo(h1))
|
||||||
assert.Equal(t, h1.localIndexId, prim.localIndexId)
|
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a))
|
||||||
assert.Equal(t, h2.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
|
|
||||||
assert.Nil(t, h5.next)
|
|
||||||
|
|
||||||
// Delete primary
|
// Deleting the same hostinfo again must not report final while siblings remain and must not
|
||||||
hm.DeleteHostInfo(h1)
|
// disturb the list. The old chain code got this wrong: the first delete nil'd next/prev, so a
|
||||||
assert.Nil(t, h1.prev)
|
// second delete looked final and wiped lighthouse state out from under the live sibling.
|
||||||
assert.Nil(t, h1.next)
|
assert.False(t, hm.DeleteHostInfo(h1))
|
||||||
|
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a))
|
||||||
|
|
||||||
// Make sure we go h2 -> h3 -> h4 -> h5
|
// Delete a middle node.
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
assert.False(t, hm.DeleteHostInfo(h3))
|
||||||
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
assert.Equal(t, []uint32{2, 4, 5}, chainIds(t, hm, a))
|
||||||
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
|
|
||||||
assert.Nil(t, h5.next)
|
|
||||||
|
|
||||||
// Delete in the middle
|
// Delete the tail.
|
||||||
hm.DeleteHostInfo(h3)
|
assert.False(t, hm.DeleteHostInfo(h5))
|
||||||
assert.Nil(t, h3.prev)
|
assert.Equal(t, []uint32{2, 4}, chainIds(t, hm, a))
|
||||||
assert.Nil(t, h3.next)
|
|
||||||
|
|
||||||
// Make sure we go h2 -> h4 -> h5
|
// Delete the head; h4 remains and becomes primary.
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
assert.False(t, hm.DeleteHostInfo(h2))
|
||||||
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
assert.Equal(t, []uint32{4}, chainIds(t, hm, a))
|
||||||
assert.Equal(t, h4.localIndexId, prim.next.localIndexId)
|
assert.Equal(t, h4, hm.QueryVpnAddr(a))
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
|
|
||||||
assert.Nil(t, h5.next)
|
|
||||||
|
|
||||||
// Delete the tail
|
// Delete the only remaining item; final is true and the address is gone.
|
||||||
hm.DeleteHostInfo(h5)
|
assert.True(t, hm.DeleteHostInfo(h4))
|
||||||
assert.Nil(t, h5.prev)
|
assert.Empty(t, chainIds(t, hm, a))
|
||||||
assert.Nil(t, h5.next)
|
assert.Nil(t, hm.QueryVpnAddr(a))
|
||||||
|
|
||||||
// Make sure we go h2 -> h4
|
// Deleting an already-gone hostinfo is still final; nothing holds the address anymore.
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
assert.True(t, hm.DeleteHostInfo(h4))
|
||||||
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
assert.Empty(t, chainIds(t, hm, a))
|
||||||
assert.Equal(t, h4.localIndexId, prim.next.localIndexId)
|
}
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Nil(t, h4.next)
|
|
||||||
|
|
||||||
// Delete the head
|
// TestHostMap_MakePrimary_DeletedHostInfo covers promoting a hostinfo that lost a race with
|
||||||
hm.DeleteHostInfo(h2)
|
// tunnel teardown: swapPrimary and AddRelay decide to promote while holding a stale pointer and
|
||||||
assert.Nil(t, h2.prev)
|
// only take the write lock after a delete fully unlinked the hostinfo. MakePrimary must be a
|
||||||
assert.Nil(t, h2.next)
|
// no-op, not a resurrection that installs an unmanaged primary.
|
||||||
|
func TestHostMap_MakePrimary_DeletedHostInfo(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
hm := newHostMap(l)
|
||||||
|
f := &Interface{}
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
|
||||||
// Make sure we only have h4
|
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
||||||
assert.Equal(t, h4.localIndexId, prim.localIndexId)
|
hm.unlockedAddHostInfo(h1, f)
|
||||||
assert.Nil(t, prim.prev)
|
hm.unlockedAddHostInfo(h2, f)
|
||||||
assert.Nil(t, prim.next)
|
|
||||||
assert.Nil(t, h4.next)
|
|
||||||
|
|
||||||
// Delete the only item
|
// h1 is fully deleted while another goroutine still holds a pointer to it.
|
||||||
hm.DeleteHostInfo(h4)
|
assert.False(t, hm.DeleteHostInfo(h1))
|
||||||
assert.Nil(t, h4.prev)
|
assert.Equal(t, []uint32{2}, chainIds(t, hm, a))
|
||||||
assert.Nil(t, h4.next)
|
|
||||||
|
|
||||||
// Make sure we have nil
|
// The stale promote must not bring it back.
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
hm.MakePrimary(h1)
|
||||||
assert.Nil(t, prim)
|
assert.Equal(t, []uint32{2}, chainIds(t, hm, a))
|
||||||
|
assert.Equal(t, h2, hm.QueryVpnAddr(a))
|
||||||
|
assert.Nil(t, hm.QueryIndex(h1.localIndexId))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestHostMap_QueryVpnAddrsRelayFor_NonPrimary makes sure a relay established on an older
|
||||||
|
// hostinfo is still found after a newer tunnel without relay state takes primary for the same
|
||||||
|
// address. The lookup checks the primary first and falls back to the rest of the list.
|
||||||
|
func TestHostMap_QueryVpnAddrsRelayFor_NonPrimary(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
hm := newHostMap(l)
|
||||||
|
f := &Interface{}
|
||||||
|
relayAddr := netip.MustParseAddr("0.0.0.9")
|
||||||
|
target := netip.MustParseAddr("0.0.0.1")
|
||||||
|
|
||||||
|
older := &HostInfo{
|
||||||
|
vpnAddrs: []netip.Addr{relayAddr},
|
||||||
|
localIndexId: 1,
|
||||||
|
relayState: RelayState{
|
||||||
|
relayForByAddr: map[netip.Addr]*Relay{},
|
||||||
|
relayForByIdx: map[uint32]*Relay{},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
older.relayState.InsertRelay(target, 100, &Relay{Type: ForwardingType, State: Established, LocalIndex: 100, PeerAddr: target})
|
||||||
|
hm.unlockedAddHostInfo(older, f)
|
||||||
|
|
||||||
|
// The relay is found on the primary.
|
||||||
|
h, r, err := hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, older, h)
|
||||||
|
assert.Equal(t, uint32(100), r.LocalIndex)
|
||||||
|
|
||||||
|
// A re-handshake with no relay state takes primary; the established relay on the older
|
||||||
|
// hostinfo must still be found through the fallback.
|
||||||
|
newer := &HostInfo{vpnAddrs: []netip.Addr{relayAddr}, localIndexId: 2}
|
||||||
|
hm.unlockedAddHostInfo(newer, f)
|
||||||
|
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, relayAddr))
|
||||||
|
|
||||||
|
h, r, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, older, h)
|
||||||
|
assert.Equal(t, uint32(100), r.LocalIndex)
|
||||||
|
|
||||||
|
// No hostinfo at all is a plain miss.
|
||||||
|
_, _, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, netip.MustParseAddr("0.0.0.42"))
|
||||||
|
require.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestHostMap_DeleteHostInfo_MultipleVpnAddrs exercises the case where a hostinfo carries more than one
|
||||||
|
// vpnAddr and shares its next/prev chain with a live sibling. Deleting the head must not corrupt the
|
||||||
|
// sibling: every address the sibling owns has to keep pointing at it. The pre-fix code unlinked the shared
|
||||||
|
// chain once per vpnAddr, so on the first address it nil'd next/prev, and on the second address the node
|
||||||
|
// looked already-detached: it dropped the map entry instead of promoting the sibling (and tripped the
|
||||||
|
// isLastHostinfo relay teardown). See unlockedDeleteHostInfo.
|
||||||
|
func TestHostMap_DeleteHostInfo_MultipleVpnAddrs(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
hm := newHostMap(l)
|
||||||
|
|
||||||
|
f := &Interface{}
|
||||||
|
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
b := netip.MustParseAddr("0.0.0.2")
|
||||||
|
|
||||||
|
// Two tunnels for the same peer, each reachable at both a and b.
|
||||||
|
other := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 1}
|
||||||
|
head := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 2}
|
||||||
|
|
||||||
|
hm.unlockedAddHostInfo(other, f)
|
||||||
|
hm.unlockedAddHostInfo(head, f)
|
||||||
|
|
||||||
|
// head is primary for both addresses, other is next in each address's list.
|
||||||
|
assert.Equal(t, head, hm.QueryVpnAddr(a))
|
||||||
|
assert.Equal(t, head, hm.QueryVpnAddr(b))
|
||||||
|
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a))
|
||||||
|
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, b))
|
||||||
|
|
||||||
|
// Delete the head. other is still live, so it must become primary for BOTH addresses.
|
||||||
|
assert.False(t, hm.DeleteHostInfo(head))
|
||||||
|
assert.Equal(t, other, hm.QueryVpnAddr(a))
|
||||||
|
assert.Equal(t, other, hm.QueryVpnAddr(b))
|
||||||
|
assert.Equal(t, []uint32{1}, chainIds(t, hm, a))
|
||||||
|
assert.Equal(t, []uint32{1}, chainIds(t, hm, b))
|
||||||
|
|
||||||
|
// head is fully removed from the index map.
|
||||||
|
assert.Nil(t, hm.QueryIndex(head.localIndexId))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestHostMap_DeleteHostInfo_DivergentVpnAddrs covers chained hostinfos for the same peer whose
|
||||||
|
// vpnAddrs sets differ (a re-handshake cert added a second address). Deleting the superset node
|
||||||
|
// must not promote a sibling to an address it does not own.
|
||||||
|
func TestHostMap_DeleteHostInfo_DivergentVpnAddrs(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
hm := newHostMap(l)
|
||||||
|
f := &Interface{}
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
b := netip.MustParseAddr("0.0.0.2")
|
||||||
|
|
||||||
|
// sub owns only a; super (a newer handshake) owns a and b.
|
||||||
|
sub := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||||
|
super := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 2}
|
||||||
|
hm.unlockedAddHostInfo(sub, f)
|
||||||
|
hm.unlockedAddHostInfo(super, f)
|
||||||
|
|
||||||
|
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a))
|
||||||
|
assert.Equal(t, []uint32{2}, chainIds(t, hm, b))
|
||||||
|
|
||||||
|
// Delete super: a promotes to sub (which owns it); b has no remaining owner and must be
|
||||||
|
// removed, not dangled at sub (which does not own b).
|
||||||
|
assert.False(t, hm.DeleteHostInfo(super))
|
||||||
|
assert.Equal(t, []uint32{1}, chainIds(t, hm, a))
|
||||||
|
assert.Empty(t, chainIds(t, hm, b))
|
||||||
|
assert.Equal(t, sub, hm.QueryVpnAddr(a))
|
||||||
|
assert.Nil(t, hm.QueryVpnAddr(b))
|
||||||
|
assert.Nil(t, hm.QueryIndex(super.localIndexId))
|
||||||
|
|
||||||
|
// Deleting sub cleans up fully.
|
||||||
|
assert.True(t, hm.DeleteHostInfo(sub))
|
||||||
|
assert.Nil(t, hm.QueryVpnAddr(a))
|
||||||
|
assertHostMapInvariants(t, hm)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestHostMap_AddDivergentOverlap covers a new hostinfo claiming addresses currently owned by two
|
||||||
|
// DIFFERENT hostinfos. The old single shared next/prev chain overwrote a pointer and orphaned one
|
||||||
|
// of them (in Indexes but unreachable via its address); independent per-address lists cannot.
|
||||||
|
func TestHostMap_AddDivergentOverlap(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
hm := newHostMap(l)
|
||||||
|
f := &Interface{}
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
b := netip.MustParseAddr("0.0.0.2")
|
||||||
|
|
||||||
|
hiA := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||||
|
hiP := &HostInfo{vpnAddrs: []netip.Addr{b}, localIndexId: 2}
|
||||||
|
hm.unlockedAddHostInfo(hiA, f)
|
||||||
|
hm.unlockedAddHostInfo(hiP, f)
|
||||||
|
|
||||||
|
hiB := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 3}
|
||||||
|
hm.unlockedAddHostInfo(hiB, f)
|
||||||
|
|
||||||
|
assert.Equal(t, []uint32{3, 1}, chainIds(t, hm, a))
|
||||||
|
assert.Equal(t, []uint32{3, 2}, chainIds(t, hm, b))
|
||||||
|
// hiA is still reachable via its address (not orphaned) and still indexed.
|
||||||
|
assert.Contains(t, chainIds(t, hm, a), hiA.localIndexId)
|
||||||
|
assert.NotNil(t, hm.QueryIndex(hiA.localIndexId))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs verifies the MaxHostInfosPerVpnIp overflow prune
|
||||||
|
// (unlockedInnerAddHostInfo calls unlockedDeleteHostInfo on the oldest node once the chain is too long)
|
||||||
|
// still behaves when hostinfos carry more than one vpnAddr. The pruned node is always the tail, so it is
|
||||||
|
// primary for none of the addresses, and both address chains must stay consistent afterwards.
|
||||||
|
func TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
hm := newHostMap(l)
|
||||||
|
|
||||||
|
f := &Interface{}
|
||||||
|
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
b := netip.MustParseAddr("0.0.0.2")
|
||||||
|
|
||||||
|
// Add one more than the cap, newest last so it becomes head. Every hostinfo owns both a and b.
|
||||||
|
hostinfos := make([]*HostInfo, 0, MaxHostInfosPerVpnIp+1)
|
||||||
|
for i := 0; i <= MaxHostInfosPerVpnIp; i++ {
|
||||||
|
hostinfos = append(hostinfos, &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: uint32(i + 1)})
|
||||||
|
}
|
||||||
|
// Add oldest first (highest index in our slice) so the very first one added is the overflow victim.
|
||||||
|
for i := len(hostinfos) - 1; i >= 0; i-- {
|
||||||
|
hm.unlockedAddHostInfo(hostinfos[i], f)
|
||||||
|
}
|
||||||
|
|
||||||
|
oldest := hostinfos[len(hostinfos)-1]
|
||||||
|
|
||||||
|
// The oldest hostinfo was pruned from both lists and the index map.
|
||||||
|
assert.Nil(t, hm.QueryIndex(oldest.localIndexId))
|
||||||
|
|
||||||
|
// Both addresses hold exactly MaxHostInfosPerVpnIp survivors in the same order; oldest is absent.
|
||||||
|
require.Len(t, chainIds(t, hm, a), MaxHostInfosPerVpnIp)
|
||||||
|
assert.Equal(t, chainIds(t, hm, a), chainIds(t, hm, b), "both addresses must list the same survivors in the same order")
|
||||||
|
assert.NotContains(t, chainIds(t, hm, a), oldest.localIndexId)
|
||||||
|
assert.Equal(t, hm.QueryVpnAddr(a), hm.QueryVpnAddr(b))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHostMap_reload(t *testing.T) {
|
func TestHostMap_reload(t *testing.T) {
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
@@ -9,10 +10,24 @@ import (
|
|||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
|
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
|
||||||
|
// only valid until the next Read on that queue. Every consumer below
|
||||||
|
// (parse, self-forward, handshake cache, sendInsideMessage) reads it
|
||||||
|
// synchronously; do not retain pkt outside this call. If a future
|
||||||
|
// caller needs to keep the packet, use pkt.Clone() to detach it from
|
||||||
|
// the borrow.
|
||||||
|
//
|
||||||
|
// pkt.Bytes is either one IP datagram (GSO zero) or a TSO/USO
|
||||||
|
// superpacket. In both cases the L3+L4 headers at the start describe
|
||||||
|
// the same 5-tuple every segment will share, so a single newPacket /
|
||||||
|
// firewall check covers the whole superpacket.
|
||||||
|
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 +52,14 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
||||||
// TUN device.
|
// TUN device.
|
||||||
if immediatelyForwardToSelf {
|
if immediatelyForwardToSelf {
|
||||||
_, err := f.readers[q].Write(packet)
|
// Write copies into the kernel queue synchronously, so seg's lifetime ends at return.
|
||||||
|
// A self-forwarded superpacket would be re-handed to the
|
||||||
|
// kernel as one giant blob; segment first so the loopback
|
||||||
|
// path sees one IP datagram per Write.
|
||||||
|
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||||
|
_, werr := f.queues[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 +75,23 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
|
|
||||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
||||||
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
|
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
||||||
|
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
||||||
|
// so retaining segments past the loop is safe.
|
||||||
|
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||||
|
hh.cachePacket(f.l, header.Message, 0, seg, f.sendMessageNow, f.cachedPacketMetrics)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil && f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("Failed to segment superpacket for handshake cache",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
|
)
|
||||||
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
||||||
"vpnAddr", fwPacket.RemoteAddr,
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
@@ -73,10 +107,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, rejectBuf, q)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
||||||
"fwPacket", fwPacket,
|
"fwPacket", fwPacket,
|
||||||
@@ -86,8 +119,152 @@ 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.GetRemote(),
|
||||||
|
"counter", c,
|
||||||
|
)
|
||||||
|
// Skip this segment; the rest of the superpacket can still
|
||||||
|
// go out — TCP will retransmit anything we drop here.
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendInsideMessage encrypts a firewall-approved inside packet (or every
|
||||||
|
// segment of a TSO/USO superpacket) into the caller's batch slot for
|
||||||
|
// later sendmmsg flush. Segmentation is fused with encryption here so the
|
||||||
|
// kernel-supplied superpacket bytes never get written into a separate
|
||||||
|
// scratch arena: SegmentSuperpacket builds each segment's plaintext in
|
||||||
|
// segScratch[:segLen] in turn, and we encrypt directly into a fresh
|
||||||
|
// SendBatch slot.
|
||||||
|
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int) {
|
||||||
|
ci := hostinfo.ConnectionState
|
||||||
|
if ci.eKey == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
remote := hostinfo.GetRemote()
|
||||||
|
ecnEnabled := f.ecnEnabled.Load()
|
||||||
|
if hostinfo.lastRebindCount != f.rebindCount {
|
||||||
|
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
||||||
|
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||||
|
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
|
||||||
|
hostinfo.lastRebindCount = f.rebindCount
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind counter",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !remote.IsValid() { //the relay path
|
||||||
|
//first, find our relay hostinfo:
|
||||||
|
var relayHostInfo *HostInfo
|
||||||
|
var relay *Relay
|
||||||
|
var err error
|
||||||
|
for _, relayIP := range hostinfo.relayState.CopyRelayIps() {
|
||||||
|
relayHostInfo, relay, err = f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP)
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.relayState.DeleteRelay(relayIP)
|
||||||
|
hostinfo.logger(f.l).Info("sendNoMetrics failed to find HostInfo",
|
||||||
|
"relay", relayIP,
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if relayHostInfo == nil || relay == nil {
|
||||||
|
//failure already logged
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
err = tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||||
|
//relay header + header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305) + relay tag
|
||||||
|
scratch := sendBatch.Reserve(header.Len + header.Len + len(seg) + 16 + 16)
|
||||||
|
|
||||||
|
innerPacket := f.sendInsideEncrypt(hostinfo, ci, seg, scratch[header.Len:], nb)
|
||||||
|
if innerPacket == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
//now we need to do a relay-encrypt:
|
||||||
|
toSend, err := f.prepareSendVia(relayHostInfo, relay, innerPacket, nb, scratch, true)
|
||||||
|
if err != nil {
|
||||||
|
//already logged
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var ecn byte
|
||||||
|
if ecnEnabled {
|
||||||
|
ecn = innerECN(seg)
|
||||||
|
}
|
||||||
|
sendBatch.Commit(toSend, relayHostInfo.GetRemote(), ecn)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to segment superpacket for relay send", "error", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||||
|
// header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305)
|
||||||
|
scratch := sendBatch.Reserve(header.Len + len(seg) + 16)
|
||||||
|
|
||||||
|
out := f.sendInsideEncrypt(hostinfo, ci, seg, scratch, nb)
|
||||||
|
if out == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var ecn byte
|
||||||
|
if ecnEnabled {
|
||||||
|
ecn = innerECN(seg)
|
||||||
|
}
|
||||||
|
sendBatch.Commit(out, remote, ecn)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to segment superpacket for send",
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// innerECN returns the 2-bit IP-level ECN codepoint of an inner IPv4 or IPv6
|
||||||
|
// packet, or 0 if pkt is too short or its IP version is unrecognized. Used at
|
||||||
|
// encap to copy the inner codepoint onto the outer carrier per RFC 6040.
|
||||||
|
func innerECN(pkt []byte) byte {
|
||||||
|
if len(pkt) < 2 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
switch pkt[0] >> 4 {
|
||||||
|
case 4:
|
||||||
|
return pkt[1] & 0x03
|
||||||
|
case 6:
|
||||||
|
return (pkt[1] >> 4) & 0x03
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||||
if !f.firewall.InSendReject {
|
if !f.firewall.OutboundSendReject {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -96,14 +273,14 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err := f.readers[q].Write(out)
|
_, err := f.queues[q].Write(out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to write to tun", "error", err)
|
f.l.Error("Failed to write to tun", "error", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
|
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
|
||||||
if !f.firewall.OutSendReject {
|
if !f.firewall.InboundSendReject {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -275,21 +452,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 +480,7 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
"headerLen", len(out),
|
"headerLen", len(out),
|
||||||
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
||||||
)
|
)
|
||||||
return
|
return nil, io.ErrShortBuffer
|
||||||
}
|
}
|
||||||
|
|
||||||
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
||||||
@@ -331,20 +500,44 @@ 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 {
|
||||||
|
// already logged by prepareSendVia
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
err = f.writers[0].WriteTo(out, via.remote)
|
|
||||||
|
err = f.writers[0].WriteTo(toSend, via.GetRemote())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
||||||
}
|
}
|
||||||
f.connectionManager.RelayUsed(relay.LocalIndex)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
||||||
if ci.eKey == nil {
|
if ci.eKey == nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
useRelay := !remote.IsValid() && !hostinfo.remote.IsValid()
|
useRelay := !remote.IsValid() && !hostinfo.GetRemote().IsValid()
|
||||||
fullOut := out
|
fullOut := out
|
||||||
|
|
||||||
if useRelay {
|
if useRelay {
|
||||||
@@ -391,7 +584,6 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
|||||||
"error", err,
|
"error", err,
|
||||||
"udpAddr", remote,
|
"udpAddr", remote,
|
||||||
"counter", c,
|
"counter", c,
|
||||||
"attemptedCounter", c,
|
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -404,8 +596,8 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
|||||||
"udpAddr", remote,
|
"udpAddr", remote,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
} else if hostinfo.remote.IsValid() {
|
} else if hr := hostinfo.GetRemote(); hr.IsValid() {
|
||||||
err = f.writers[q].WriteTo(out, hostinfo.remote)
|
err = f.writers[q].WriteTo(out, hr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
||||||
"error", err,
|
"error", err,
|
||||||
|
|||||||
+207
-62
@@ -4,20 +4,25 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"runtime"
|
||||||
|
"slices"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -47,7 +52,19 @@ type InterfaceConfig struct {
|
|||||||
reQueryWait time.Duration
|
reQueryWait time.Duration
|
||||||
|
|
||||||
ConntrackCacheTimeout time.Duration
|
ConntrackCacheTimeout time.Duration
|
||||||
l *slog.Logger
|
|
||||||
|
// CpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
|
||||||
|
// should pin to. Queue i pins to CpuAffinity[i % len(CpuAffinity)] —
|
||||||
|
// shorter lists than `routines` cycle. Empty list keeps the default
|
||||||
|
// pin-to-(i % NumCPU) behavior. Only consulted when PinThreads is true.
|
||||||
|
CpuAffinity []int
|
||||||
|
// PinThreads controls whether each TUN reader OS thread is pinned to a
|
||||||
|
// single CPU (via tun.pin_threads, default true). Pinning keeps each
|
||||||
|
// goroutine's sendmmsg on one XPS-selected NIC TX ring so per-flow
|
||||||
|
// packets stay ordered on the wire.
|
||||||
|
PinThreads bool
|
||||||
|
|
||||||
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
type Interface struct {
|
type Interface struct {
|
||||||
@@ -71,7 +88,21 @@ type Interface struct {
|
|||||||
routines int
|
routines int
|
||||||
disconnectInvalid atomic.Bool
|
disconnectInvalid atomic.Bool
|
||||||
closed atomic.Bool
|
closed atomic.Bool
|
||||||
relayManager *relayManager
|
// cpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
|
||||||
|
// should pin to. Queue i pins to cpuAffinity[i % len(cpuAffinity)].
|
||||||
|
// Empty falls back to the default pin-to-(allowed CPU) behavior.
|
||||||
|
// Only consulted when pinThreads is true.
|
||||||
|
cpuAffinity []int
|
||||||
|
// pinThreads controls whether listenIn pins each TUN reader OS thread to
|
||||||
|
// a CPU at all (tun.pin_threads, default true). When false, threads are
|
||||||
|
// left free to migrate as on stock nebula.
|
||||||
|
pinThreads bool
|
||||||
|
// ecnEnabled gates RFC 6040 underlay ECN propagation. When true,
|
||||||
|
// inside.go copies the inner ECN onto the outer carrier on encap and
|
||||||
|
// decryptToTun folds outer CE into the inner header on decap. Toggle
|
||||||
|
// via tunnels.ecn (default true).
|
||||||
|
ecnEnabled atomic.Bool
|
||||||
|
relayManager *relayManager
|
||||||
|
|
||||||
tryPromoteEvery atomic.Uint32
|
tryPromoteEvery atomic.Uint32
|
||||||
reQueryEvery atomic.Uint32
|
reQueryEvery atomic.Uint32
|
||||||
@@ -88,8 +119,12 @@ type Interface struct {
|
|||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
writers []udp.Conn
|
writers []udp.Conn
|
||||||
readers []io.ReadWriteCloser
|
queues []tio.Queue
|
||||||
wg sync.WaitGroup
|
// batchers is one per tun queue, wrapping queues[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)
|
||||||
@@ -187,7 +222,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
routines: c.routines,
|
routines: c.routines,
|
||||||
version: c.version,
|
version: c.version,
|
||||||
writers: make([]udp.Conn, c.routines),
|
writers: make([]udp.Conn, c.routines),
|
||||||
readers: make([]io.ReadWriteCloser, 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,
|
||||||
@@ -196,6 +231,8 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
relayManager: c.relayManager,
|
relayManager: c.relayManager,
|
||||||
connectionManager: c.connectionManager,
|
connectionManager: c.connectionManager,
|
||||||
conntrackCacheTimeout: c.ConntrackCacheTimeout,
|
conntrackCacheTimeout: c.ConntrackCacheTimeout,
|
||||||
|
cpuAffinity: c.CpuAffinity,
|
||||||
|
pinThreads: c.PinThreads,
|
||||||
|
|
||||||
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
||||||
messageMetrics: c.MessageMetrics,
|
messageMetrics: c.MessageMetrics,
|
||||||
@@ -213,6 +250,9 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
|
|
||||||
ifce.connectionManager.intf = ifce
|
ifce.connectionManager.intf = ifce
|
||||||
|
|
||||||
|
// Held until Close so waiting on the interface blocks until the resources are actually released
|
||||||
|
ifce.wg.Add(1)
|
||||||
|
|
||||||
return ifce, nil
|
return ifce, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -235,38 +275,52 @@ func (f *Interface) activate() error {
|
|||||||
"boringcrypto", boringEnabled(),
|
"boringcrypto", boringEnabled(),
|
||||||
)
|
)
|
||||||
|
|
||||||
if f.routines > 1 {
|
if f.routines > 1 && !f.outside.SupportsMultipleReaders() {
|
||||||
if !f.inside.SupportsMultiqueue() || !f.outside.SupportsMultipleReaders() {
|
f.routines = 1
|
||||||
f.routines = 1
|
f.l.Warn("multiple udp readers are not supported on this platform, falling back to a single routine")
|
||||||
f.l.Warn("routines is not supported on this platform, falling back to a single routine")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Prepare the tun queues. A device that can't open that many hands back
|
||||||
|
// fewer (a single queue on platforms without multiqueue support) and we
|
||||||
|
// size the reader routines to what we actually got.
|
||||||
|
queues, err := f.inside.Queues(f.routines)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if len(queues) < f.routines {
|
||||||
|
f.l.Warn("tun multiqueue is not supported on this platform, falling back to fewer routines",
|
||||||
|
"requested", f.routines, "opened", len(queues))
|
||||||
|
f.routines = len(queues)
|
||||||
|
}
|
||||||
|
f.queues = queues
|
||||||
|
|
||||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||||
|
|
||||||
// Prepare n tun queues
|
for i := range f.queues {
|
||||||
var reader io.ReadWriteCloser = f.inside
|
caps := tio.QueueCapabilities(f.queues[i])
|
||||||
for i := 0; i < f.routines; i++ {
|
if caps.TSO || caps.USO {
|
||||||
if i > 0 {
|
// Multi-lane: TCP gets coalesced when TSO is on, UDP when USO
|
||||||
reader, err = f.inside.NewMultiQueueReader()
|
// is on, everything else (and either lane disabled) falls
|
||||||
if err != nil {
|
// through to passthrough so non-IP / non-TCP-UDP traffic still
|
||||||
return err
|
// reaches the TUN.
|
||||||
}
|
arena := batch.NewArena(batch.DefaultMultiArenaCap)
|
||||||
|
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l, arena, caps.TSO, caps.USO)
|
||||||
|
} else {
|
||||||
|
arena := batch.NewArena(batch.DefaultPassthroughArenaCap)
|
||||||
|
f.batchers[i] = batch.NewPassthrough(f.queues[i], arena.Reserve, arena.Reset)
|
||||||
}
|
}
|
||||||
f.readers[i] = reader
|
|
||||||
}
|
}
|
||||||
|
|
||||||
f.wg.Add(1) // for us to wait on Close() to return
|
// On error the caller owns the cleanup, Control.Start cancels the service context
|
||||||
|
// before releasing our resources so a waiter never observes a live context
|
||||||
if err = f.inside.Activate(); err != nil {
|
if err = f.inside.Activate(); err != nil {
|
||||||
f.wg.Done()
|
|
||||||
f.inside.Close()
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) run() (func() error, error) {
|
func (f *Interface) run() {
|
||||||
// Launch n queues to read packets from udp
|
// Launch n queues to read packets from udp
|
||||||
for i := 0; i < f.routines; i++ {
|
for i := 0; i < f.routines; i++ {
|
||||||
f.wg.Go(func() {
|
f.wg.Go(func() {
|
||||||
@@ -277,17 +331,18 @@ func (f *Interface) run() (func() error, error) {
|
|||||||
// Launch n queues to read packets from tun dev
|
// Launch n queues to read packets from tun dev
|
||||||
for i := 0; i < f.routines; i++ {
|
for i := 0; i < f.routines; i++ {
|
||||||
f.wg.Go(func() {
|
f.wg.Go(func() {
|
||||||
f.listenIn(f.readers[i], i)
|
f.listenIn(f.queues[i], i)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
return func() error {
|
}
|
||||||
f.wg.Wait()
|
|
||||||
if e := f.fatalErr.Load(); e != nil {
|
func (f *Interface) wait() error {
|
||||||
return *e
|
f.wg.Wait()
|
||||||
}
|
if e := f.fatalErr.Load(); e != nil {
|
||||||
return nil
|
return *e
|
||||||
}, nil
|
}
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one
|
// onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one
|
||||||
@@ -311,16 +366,27 @@ 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())
|
plaintext := f.batchers[i].Reserve(len(payload))
|
||||||
})
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(), meta)
|
||||||
|
}
|
||||||
|
|
||||||
if err != nil && !f.closed.Load() {
|
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)
|
||||||
|
|
||||||
|
// An error after teardown began is shutdown noise, the closed flag covers resources
|
||||||
|
// Close releases itself and the cancelled ctx covers ones torn down by their owners
|
||||||
|
// reacting to it, like the user device pipes
|
||||||
|
if err != nil && !f.closed.Load() && f.ctx.Err() == nil {
|
||||||
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
||||||
f.onFatal(err)
|
f.onFatal(err)
|
||||||
}
|
}
|
||||||
@@ -328,25 +394,61 @@ 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(queue tio.Queue, i int) {
|
||||||
packet := make([]byte, mtu)
|
// Pinning this thread (and goroutine) to a single CPU keeps every sendmmsg from this goroutine going through the
|
||||||
out := make([]byte, mtu)
|
// same TX ring on the nic, so the wire sees per-flow order. Skip entirely when tun.pin_threads is false.
|
||||||
|
if f.pinThreads {
|
||||||
|
var cpu int
|
||||||
|
if n := len(f.cpuAffinity); n > 0 {
|
||||||
|
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already
|
||||||
|
// validated the entries against the allowed CPU set.
|
||||||
|
cpu = f.cpuAffinity[i%n]
|
||||||
|
} else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 {
|
||||||
|
// Default: spread queues across the CPUs we're actually allowed to
|
||||||
|
// run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so
|
||||||
|
// i % NumCPU would pick unrunnable IDs and every pin would fail.
|
||||||
|
cpu = allowed[i%len(allowed)]
|
||||||
|
} else {
|
||||||
|
cpu = i % runtime.NumCPU()
|
||||||
|
}
|
||||||
|
if err := util.PinThreadToCPU(cpu); err != nil {
|
||||||
|
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
rejectBuf := make([]byte, mtu)
|
||||||
|
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
|
||||||
|
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
n, err := reader.Read(packet)
|
pkts, err := queue.Read()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if !f.closed.Load() {
|
// Same shutdown noise handling as listenOut
|
||||||
|
if !f.closed.Load() && f.ctx.Err() == nil {
|
||||||
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
||||||
f.onFatal(err)
|
f.onFatal(err)
|
||||||
}
|
}
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get())
|
for _, pkt := range pkts {
|
||||||
|
f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get())
|
||||||
|
// Flush incrementally once a full sendmmsg batch has
|
||||||
|
// accumulated so the first packets of a deep read drain
|
||||||
|
// hit the wire while the rest are still being encrypted.
|
||||||
|
if sb.Len() >= batch.SendBatchCap {
|
||||||
|
if err := sb.Flush(); err != nil {
|
||||||
|
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := sb.Flush(); err != nil {
|
||||||
|
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
f.l.Debug("overlay reader is done", "reader", i)
|
f.l.Debug("overlay reader is done", "reader", i)
|
||||||
@@ -358,6 +460,7 @@ func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
|||||||
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
||||||
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
||||||
c.RegisterReloadCallback(f.reloadMisc)
|
c.RegisterReloadCallback(f.reloadMisc)
|
||||||
|
c.RegisterReloadCallback(f.reloadEcn)
|
||||||
|
|
||||||
for _, udpConn := range f.writers {
|
for _, udpConn := range f.writers {
|
||||||
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
||||||
@@ -375,13 +478,22 @@ func (f *Interface) reloadDisconnectInvalid(c *config.C) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) reloadFirewall(c *config.C) {
|
func (f *Interface) reloadFirewall(c *config.C) {
|
||||||
//TODO: need to trigger/detect if the certificate changed too
|
cs := f.pki.getCertState()
|
||||||
if c.HasChanged("firewall") == false {
|
curCert := cs.getCertificate(cert.Version2)
|
||||||
|
if curCert == nil {
|
||||||
|
curCert = cs.getCertificate(cert.Version1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The firewall builds its routableNetworks set from the certificate's UnsafeNetworks at construction.
|
||||||
|
// Check to see if that set has changed, and if so, rebuild the firewall.
|
||||||
|
certUnsafeChanged := curCert != nil && !slices.Equal(curCert.UnsafeNetworks(), f.firewall.unsafeNetworks)
|
||||||
|
|
||||||
|
if !c.HasChanged("firewall") && !certUnsafeChanged {
|
||||||
f.l.Debug("No firewall config change detected")
|
f.l.Debug("No firewall config change detected")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
fw, err := NewFirewallFromConfig(f.l, f.pki.getCertState(), c)
|
fw, err := NewFirewallFromConfig(f.l, cs, c)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Error while creating firewall during reload", "error", err)
|
f.l.Error("Error while creating firewall during reload", "error", err)
|
||||||
return
|
return
|
||||||
@@ -481,6 +593,23 @@ func (f *Interface) reloadMisc(c *config.C) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// reloadEcn syncs Interface.ecnEnabled with the tunnels.ecn config knob.
|
||||||
|
// Default is enabled (RFC 6040 normal mode); set false on the rare path
|
||||||
|
// where an underlay middlebox rewrites or drops ECN bits unpredictably.
|
||||||
|
func (f *Interface) reloadEcn(c *config.C) {
|
||||||
|
initial := c.InitialLoad()
|
||||||
|
if initial || c.HasChanged("tunnels.ecn") {
|
||||||
|
v := c.GetBool("tunnels.ecn", true)
|
||||||
|
changed := f.ecnEnabled.Swap(v) != v
|
||||||
|
if !initial {
|
||||||
|
f.l.Info("tunnels.ecn changed", "enabled", v)
|
||||||
|
if changed {
|
||||||
|
f.l.Warn("tunnels.ecn datapath toggled, but route-level ECN negotiation (RTAX_FEATURE_ECN) retains its previous state until nebula is restarted", "enabled", v)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
||||||
ticker := time.NewTicker(i)
|
ticker := time.NewTicker(i)
|
||||||
defer ticker.Stop()
|
defer ticker.Stop()
|
||||||
@@ -491,26 +620,34 @@ func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
|||||||
certInitiatingVersion := metrics.GetOrRegisterGauge("certificate.initiating_version", nil)
|
certInitiatingVersion := metrics.GetOrRegisterGauge("certificate.initiating_version", nil)
|
||||||
certMaxVersion := metrics.GetOrRegisterGauge("certificate.max_version", nil)
|
certMaxVersion := metrics.GetOrRegisterGauge("certificate.max_version", nil)
|
||||||
|
|
||||||
|
emit := func() {
|
||||||
|
f.firewall.EmitStats()
|
||||||
|
f.handshakeManager.EmitStats()
|
||||||
|
udpStats()
|
||||||
|
|
||||||
|
certState := f.pki.getCertState()
|
||||||
|
defaultCrt := certState.GetDefaultCertificate()
|
||||||
|
certExpirationGauge.Update(int64(defaultCrt.NotAfter().Sub(time.Now()) / time.Second))
|
||||||
|
certInitiatingVersion.Update(int64(defaultCrt.Version()))
|
||||||
|
|
||||||
|
// Report the max certificate version we are capable of using
|
||||||
|
if certState.v2Cert != nil {
|
||||||
|
certMaxVersion.Update(int64(certState.v2Cert.Version()))
|
||||||
|
} else {
|
||||||
|
certMaxVersion.Update(int64(certState.v1Cert.Version()))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prime gauges so a Prometheus scrape that lands before the first tick
|
||||||
|
// sees real values instead of the zero defaults (issue #907).
|
||||||
|
emit()
|
||||||
|
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
f.firewall.EmitStats()
|
emit()
|
||||||
f.handshakeManager.EmitStats()
|
|
||||||
udpStats()
|
|
||||||
|
|
||||||
certState := f.pki.getCertState()
|
|
||||||
defaultCrt := certState.GetDefaultCertificate()
|
|
||||||
certExpirationGauge.Update(int64(defaultCrt.NotAfter().Sub(time.Now()) / time.Second))
|
|
||||||
certInitiatingVersion.Update(int64(defaultCrt.Version()))
|
|
||||||
|
|
||||||
// Report the max certificate version we are capable of using
|
|
||||||
if certState.v2Cert != nil {
|
|
||||||
certMaxVersion.Update(int64(certState.v2Cert.Version()))
|
|
||||||
} else {
|
|
||||||
certMaxVersion.Update(int64(certState.v1Cert.Version()))
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -523,9 +660,15 @@ func (f *Interface) GetCertState() *CertState {
|
|||||||
return f.pki.getCertState()
|
return f.pki.getCertState()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Close releases the interface's resources: the udp sockets and the tun device.
|
||||||
|
// It is idempotent and safe to call at any point in the lifecycle, including on an interface that never activated,
|
||||||
|
// calls after the first return nil without doing anything.
|
||||||
func (f *Interface) Close() error {
|
func (f *Interface) Close() error {
|
||||||
|
if !f.closed.CompareAndSwap(false, true) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
var errs []error
|
var errs []error
|
||||||
f.closed.Store(true)
|
|
||||||
|
|
||||||
// Release the udp readers
|
// Release the udp readers
|
||||||
for i, u := range f.writers {
|
for i, u := range f.writers {
|
||||||
@@ -541,6 +684,8 @@ func (f *Interface) Close() error {
|
|||||||
if closeErr != nil {
|
if closeErr != nil {
|
||||||
errs = append(errs, closeErr)
|
errs = append(errs, closeErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Release the construction token so waiters know the resources are gone
|
||||||
f.wg.Done()
|
f.wg.Done()
|
||||||
return errors.Join(errs...)
|
return errors.Join(errs...)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,73 @@
|
|||||||
|
//go:build linux || darwin
|
||||||
|
|
||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/firewall"
|
||||||
|
"github.com/slackhq/nebula/overlay/overlaytest"
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Test_emitStats_primesGauges covers issue #907: a Prometheus scrape that
|
||||||
|
// landed before the first ticker fire used to read 0 for the cert gauges.
|
||||||
|
// emitStats now primes the gauges before entering the ticker loop. We assert
|
||||||
|
// the gauge is zero before the first call and non-zero after.
|
||||||
|
func Test_emitStats_primesGauges(t *testing.T) {
|
||||||
|
defer metrics.DefaultRegistry.UnregisterAll()
|
||||||
|
|
||||||
|
l := test.NewLogger()
|
||||||
|
hostMap := newHostMap(l)
|
||||||
|
preferredRanges := []netip.Prefix{netip.MustParsePrefix("10.0.0.0/8")}
|
||||||
|
hostMap.preferredRanges.Store(&preferredRanges)
|
||||||
|
|
||||||
|
notAfter := time.Now().Add(time.Hour)
|
||||||
|
cs := &CertState{
|
||||||
|
initiatingVersion: cert.Version1,
|
||||||
|
privateKey: []byte{},
|
||||||
|
v1Cert: &dummyCert{version: cert.Version1, notAfter: notAfter},
|
||||||
|
v1Credential: nil,
|
||||||
|
}
|
||||||
|
|
||||||
|
lh := newTestLighthouse()
|
||||||
|
ifce := &Interface{
|
||||||
|
hostMap: hostMap,
|
||||||
|
inside: &overlaytest.NoopTun{},
|
||||||
|
outside: &udp.NoopConn{},
|
||||||
|
firewall: &Firewall{Conntrack: &FirewallConntrack{Conns: map[firewall.Packet]*conn{}}},
|
||||||
|
lightHouse: lh,
|
||||||
|
pki: &PKI{},
|
||||||
|
handshakeManager: NewHandshakeManager(l, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
||||||
|
l: l,
|
||||||
|
// On linux, udp.NewUDPStatsEmitter indexes writers[0] and asserts to
|
||||||
|
// *udp.StdConn. A zero value works: getMemInfo sees a nil rawConn,
|
||||||
|
// returns an error, and the emitter falls through to a no-op.
|
||||||
|
writers: []udp.Conn{&udp.StdConn{}},
|
||||||
|
}
|
||||||
|
ifce.pki.cs.Store(cs)
|
||||||
|
|
||||||
|
ttlGauge := metrics.GetOrRegisterGauge("certificate.ttl_seconds", nil)
|
||||||
|
require.Zero(t, ttlGauge.Value(), "gauge should be zero before emitStats runs")
|
||||||
|
|
||||||
|
// Pre-cancel the context so emitStats returns after priming the gauges
|
||||||
|
// without ever reading from ticker.C. The one hour interval is just a
|
||||||
|
// belt-and-suspenders, the test does not expect the ticker to fire.
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
cancel()
|
||||||
|
ifce.emitStats(ctx, time.Hour)
|
||||||
|
|
||||||
|
ttl := ttlGauge.Value()
|
||||||
|
assert.Positive(t, ttl, "ttl gauge should be primed by emitStats before its first tick")
|
||||||
|
assert.LessOrEqual(t, ttl, int64(3600))
|
||||||
|
assert.Equal(t, int64(cert.Version1), metrics.GetOrRegisterGauge("certificate.initiating_version", nil).Value())
|
||||||
|
assert.Equal(t, int64(cert.Version1), metrics.GetOrRegisterGauge("certificate.max_version", nil).Value())
|
||||||
|
}
|
||||||
@@ -0,0 +1,120 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestReloadFirewall_CertUnsafeNetworksChanged verifies that reloadFirewall
|
||||||
|
// rebuilds the firewall when only the certificate's UnsafeNetworks have changed,
|
||||||
|
// even if the firewall section of the YAML has not.
|
||||||
|
func TestReloadFirewall_CertUnsafeNetworksChanged(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
|
||||||
|
vpnNet := netip.MustParsePrefix("10.0.0.1/24")
|
||||||
|
initialUnsafe := []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")}
|
||||||
|
|
||||||
|
// dummyCert avoids dragging the real signing pipeline into a unit test.
|
||||||
|
c1 := &dummyCert{
|
||||||
|
version: cert.Version2,
|
||||||
|
networks: []netip.Prefix{vpnNet},
|
||||||
|
unsafeNetworks: initialUnsafe,
|
||||||
|
}
|
||||||
|
pki := &PKI{}
|
||||||
|
pki.cs.Store(&CertState{v2Cert: c1, initiatingVersion: cert.Version2})
|
||||||
|
|
||||||
|
rawYAML := `firewall:
|
||||||
|
outbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
inbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
`
|
||||||
|
cfg := config.NewC(l)
|
||||||
|
require.NoError(t, cfg.LoadString(rawYAML))
|
||||||
|
|
||||||
|
fw, err := NewFirewallFromConfig(l, pki.getCertState(), cfg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, initialUnsafe, fw.unsafeNetworks)
|
||||||
|
|
||||||
|
f := &Interface{
|
||||||
|
pki: pki,
|
||||||
|
firewall: fw,
|
||||||
|
l: l,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Swap the cert with a different UnsafeNetworks set.
|
||||||
|
newUnsafe := []netip.Prefix{
|
||||||
|
netip.MustParsePrefix("198.51.100.0/24"),
|
||||||
|
netip.MustParsePrefix("203.0.113.0/24"),
|
||||||
|
}
|
||||||
|
c2 := &dummyCert{
|
||||||
|
version: cert.Version2,
|
||||||
|
networks: []netip.Prefix{vpnNet},
|
||||||
|
unsafeNetworks: newUnsafe,
|
||||||
|
}
|
||||||
|
pki.cs.Store(&CertState{v2Cert: c2, initiatingVersion: cert.Version2})
|
||||||
|
|
||||||
|
// Reload with the same YAML so HasChanged("firewall") reports false.
|
||||||
|
require.NoError(t, cfg.ReloadConfigString(rawYAML))
|
||||||
|
require.False(t, cfg.HasChanged("firewall"))
|
||||||
|
|
||||||
|
f.reloadFirewall(cfg)
|
||||||
|
|
||||||
|
assert.NotSame(t, fw, f.firewall, "firewall pointer should have been replaced")
|
||||||
|
assert.Equal(t, newUnsafe, f.firewall.unsafeNetworks)
|
||||||
|
assert.True(t, f.firewall.routableNetworks.Contains(netip.MustParseAddr("203.0.113.5")))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestReloadFirewall_NoChange verifies that reloadFirewall is a no-op when
|
||||||
|
// neither the firewall config nor the cert's UnsafeNetworks have changed.
|
||||||
|
func TestReloadFirewall_NoChange(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
|
||||||
|
vpnNet := netip.MustParsePrefix("10.0.0.1/24")
|
||||||
|
unsafe := []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")}
|
||||||
|
|
||||||
|
c1 := &dummyCert{
|
||||||
|
version: cert.Version2,
|
||||||
|
networks: []netip.Prefix{vpnNet},
|
||||||
|
unsafeNetworks: unsafe,
|
||||||
|
}
|
||||||
|
pki := &PKI{}
|
||||||
|
pki.cs.Store(&CertState{v2Cert: c1, initiatingVersion: cert.Version2})
|
||||||
|
|
||||||
|
rawYAML := `firewall:
|
||||||
|
outbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
inbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
`
|
||||||
|
cfg := config.NewC(l)
|
||||||
|
require.NoError(t, cfg.LoadString(rawYAML))
|
||||||
|
|
||||||
|
fw, err := NewFirewallFromConfig(l, pki.getCertState(), cfg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
f := &Interface{
|
||||||
|
pki: pki,
|
||||||
|
firewall: fw,
|
||||||
|
l: l,
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, cfg.ReloadConfigString(rawYAML))
|
||||||
|
f.reloadFirewall(cfg)
|
||||||
|
|
||||||
|
assert.Same(t, fw, f.firewall, "firewall should not have been replaced")
|
||||||
|
}
|
||||||
+290
-19
@@ -4,26 +4,54 @@ import (
|
|||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
|
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
|
"golang.org/x/net/ipv6"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// Need 96 bytes for the largest reject packet:
|
// MaxIPv4RejectPacketSize is the largest IPv4 reject packet:
|
||||||
// - 20 byte ipv4 header
|
// - 20 byte ipv4 header
|
||||||
// - 8 byte icmpv4 header
|
// - 8 byte icmpv4 header
|
||||||
// - 68 byte body (60 byte max orig ipv4 header + 8 byte orig icmpv4 header)
|
// - 68 byte body (60 byte max orig ipv4 header + 8 byte orig icmpv4 header)
|
||||||
MaxRejectPacketSize = ipv4.HeaderLen + 8 + 60 + 8
|
maxIPv4RejectPacketSize = ipv4.HeaderLen + 8 + 60 + 8
|
||||||
|
|
||||||
|
// MaxRejectPacketSize is sized for the largest possible reject packet (IPv6):
|
||||||
|
// - 40 byte ipv6 header
|
||||||
|
// - 8 byte icmpv6 header
|
||||||
|
// - up to 1000 byte body (original packet, possibly truncated. We want to stay
|
||||||
|
// under the MTU with Nebula overhead included)
|
||||||
|
maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000
|
||||||
|
|
||||||
|
MaxRejectPacketSize = maxIPv6RejectPacketSize
|
||||||
)
|
)
|
||||||
|
|
||||||
func CreateRejectPacket(packet []byte, out []byte) []byte {
|
func CreateRejectPacket(packet []byte, out []byte) []byte {
|
||||||
if len(packet) < ipv4.HeaderLen || int(packet[0]>>4) != ipv4.Version {
|
if len(packet) < 1 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
switch packet[9] {
|
version := int(packet[0] >> 4)
|
||||||
case 6: // tcp
|
switch version {
|
||||||
return ipv4CreateRejectTCPPacket(packet, out)
|
case ipv4.Version:
|
||||||
|
if len(packet) < ipv4.HeaderLen {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// Do not send reject packets for non-first fragments
|
||||||
|
if packet[6]&0x1f != 0 || packet[7] != 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
switch packet[9] {
|
||||||
|
case 6: // tcp
|
||||||
|
return ipv4CreateRejectTCPPacket(packet, out)
|
||||||
|
default:
|
||||||
|
return ipv4CreateRejectICMPPacket(packet, out)
|
||||||
|
}
|
||||||
|
case ipv6.Version:
|
||||||
|
if len(packet) < ipv6.HeaderLen {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return ipv6CreateRejectPacket(packet, out)
|
||||||
default:
|
default:
|
||||||
return ipv4CreateRejectICMPPacket(packet, out)
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -35,12 +63,17 @@ func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ICMP reply includes original header and first 8 bytes of the packet
|
// Do not generate ICMP errors in response to ICMP error packets
|
||||||
packetLen := len(packet)
|
if packet[9] == 1 && len(packet) > ihl {
|
||||||
if packetLen > ihl+8 {
|
icmpType := packet[ihl]
|
||||||
packetLen = ihl + 8
|
if icmpType == 3 || icmpType == 4 || icmpType == 5 || icmpType == 11 || icmpType == 12 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ICMP reply includes original header and first 8 bytes of the packet
|
||||||
|
packetLen := min(len(packet), ihl+8)
|
||||||
|
|
||||||
outLen := ipv4.HeaderLen + 8 + packetLen
|
outLen := ipv4.HeaderLen + 8 + packetLen
|
||||||
if outLen > cap(out) {
|
if outLen > cap(out) {
|
||||||
return nil
|
return nil
|
||||||
@@ -71,14 +104,14 @@ func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
|
|||||||
|
|
||||||
// ICMP Destination Unreachable
|
// ICMP Destination Unreachable
|
||||||
icmpOut := out[ipv4.HeaderLen:]
|
icmpOut := out[ipv4.HeaderLen:]
|
||||||
icmpOut[0] = 3 // type (Destination unreachable)
|
icmpOut[0] = 3 // type (Destination unreachable)
|
||||||
icmpOut[1] = 3 // code (Port unreachable error)
|
icmpOut[1] = 13 // code (Communication administratively prohibited)
|
||||||
icmpOut[2] = 0 // checksum
|
icmpOut[2] = 0 // checksum
|
||||||
icmpOut[3] = 0 // .
|
icmpOut[3] = 0 // .
|
||||||
icmpOut[4] = 0 // unused
|
icmpOut[4] = 0 // unused
|
||||||
icmpOut[5] = 0 // .
|
icmpOut[5] = 0 // .
|
||||||
icmpOut[6] = 0 // .
|
icmpOut[6] = 0 // .
|
||||||
icmpOut[7] = 0 // .
|
icmpOut[7] = 0 // .
|
||||||
|
|
||||||
// Copy original IP header and first 8 bytes as body
|
// Copy original IP header and first 8 bytes as body
|
||||||
copy(icmpOut[8:], packet[:packetLen])
|
copy(icmpOut[8:], packet[:packetLen])
|
||||||
@@ -165,7 +198,193 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte {
|
|||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
|
||||||
|
proto, offset, isFragment := ipv6FindUpperProtocol(packet)
|
||||||
|
if isFragment {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
switch proto {
|
||||||
|
case 6: // tcp
|
||||||
|
return ipv6CreateRejectTCPPacket(packet, out, offset)
|
||||||
|
default:
|
||||||
|
return ipv6CreateRejectICMPPacket(packet, out, proto, offset)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func ipv6CreateRejectICMPPacket(packet []byte, out []byte, proto uint8, offset int) []byte {
|
||||||
|
// Do not generate ICMPv6 errors in response to ICMPv6 error packets
|
||||||
|
if proto == 58 && len(packet) > offset {
|
||||||
|
icmpType := packet[offset]
|
||||||
|
if icmpType >= 1 && icmpType <= 4 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Include as much of the original packet as possible, up to 1000 bytes,
|
||||||
|
// so the response fits comfortably within any tunnel MTU.
|
||||||
|
packetLen := min(len(packet), 1000)
|
||||||
|
|
||||||
|
outLen := ipv6.HeaderLen + 8 + packetLen
|
||||||
|
if outLen > cap(out) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
out = out[:outLen]
|
||||||
|
|
||||||
|
// IPv6 header
|
||||||
|
ipHdr := out[0:ipv6.HeaderLen]
|
||||||
|
ipHdr[0] = ipv6.Version << 4 // version, traffic class (high bits)
|
||||||
|
ipHdr[1] = 0 // traffic class (low bits), flow label (high bits)
|
||||||
|
ipHdr[2] = 0 // flow label
|
||||||
|
ipHdr[3] = 0 // flow label
|
||||||
|
|
||||||
|
payloadLen := uint16(outLen - ipv6.HeaderLen)
|
||||||
|
binary.BigEndian.PutUint16(ipHdr[4:], payloadLen) // payload length
|
||||||
|
ipHdr[6] = 58 // next header (ICMPv6)
|
||||||
|
ipHdr[7] = 64 // hop limit
|
||||||
|
|
||||||
|
// Swap dest / src IPs (each 16 bytes, src at 8, dst at 24)
|
||||||
|
copy(ipHdr[8:24], packet[24:40])
|
||||||
|
copy(ipHdr[24:40], packet[8:24])
|
||||||
|
|
||||||
|
// ICMPv6 Destination Unreachable
|
||||||
|
icmpOut := out[ipv6.HeaderLen:]
|
||||||
|
icmpOut[0] = 1 // type (Destination Unreachable)
|
||||||
|
icmpOut[1] = 1 // code (Communication with destination administratively prohibited)
|
||||||
|
icmpOut[2] = 0 // checksum
|
||||||
|
icmpOut[3] = 0 // .
|
||||||
|
icmpOut[4] = 0 // unused
|
||||||
|
icmpOut[5] = 0 // .
|
||||||
|
icmpOut[6] = 0 // .
|
||||||
|
icmpOut[7] = 0 // .
|
||||||
|
|
||||||
|
copy(icmpOut[8:], packet[:packetLen])
|
||||||
|
|
||||||
|
// ICMPv6 checksum uses a pseudo-header
|
||||||
|
csum := ipv6PseudoheaderChecksum(ipHdr[8:24], ipHdr[24:40], 58, uint32(payloadLen))
|
||||||
|
binary.BigEndian.PutUint16(icmpOut[2:], tcpipChecksum(icmpOut, csum))
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte {
|
||||||
|
const tcpLen = 20
|
||||||
|
|
||||||
|
if len(packet) < offset+tcpLen {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
outLen := ipv6.HeaderLen + tcpLen
|
||||||
|
if outLen > cap(out) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
out = out[:outLen]
|
||||||
|
|
||||||
|
// IPv6 header
|
||||||
|
ipHdr := out[0:ipv6.HeaderLen]
|
||||||
|
ipHdr[0] = ipv6.Version << 4 // version, traffic class (high bits)
|
||||||
|
ipHdr[1] = 0 // traffic class (low bits), flow label (high bits)
|
||||||
|
ipHdr[2] = 0 // flow label
|
||||||
|
ipHdr[3] = 0 // flow label
|
||||||
|
|
||||||
|
binary.BigEndian.PutUint16(ipHdr[4:], tcpLen) // payload length
|
||||||
|
ipHdr[6] = 6 // next header (TCP)
|
||||||
|
ipHdr[7] = 64 // hop limit
|
||||||
|
|
||||||
|
// Swap dest / src IPs
|
||||||
|
copy(ipHdr[8:24], packet[24:40])
|
||||||
|
copy(ipHdr[24:40], packet[8:24])
|
||||||
|
|
||||||
|
// TCP RST
|
||||||
|
tcpIn := packet[offset:]
|
||||||
|
var ackSeq, seq uint32
|
||||||
|
outFlags := byte(0b00000100) // RST
|
||||||
|
|
||||||
|
inAck := tcpIn[13]&0b00010000 != 0
|
||||||
|
if inAck {
|
||||||
|
seq = binary.BigEndian.Uint32(tcpIn[8:])
|
||||||
|
} else {
|
||||||
|
inSyn := uint32((tcpIn[13] & 0b00000010) >> 1)
|
||||||
|
inFin := uint32(tcpIn[13] & 0b00000001)
|
||||||
|
ackSeq = binary.BigEndian.Uint32(tcpIn[4:]) + inSyn + inFin + uint32(len(tcpIn)) - uint32(tcpIn[12]>>4)<<2
|
||||||
|
outFlags |= 0b00010000 // ACK
|
||||||
|
}
|
||||||
|
|
||||||
|
tcpOut := out[ipv6.HeaderLen:]
|
||||||
|
// Swap dest / src ports
|
||||||
|
copy(tcpOut[0:2], tcpIn[2:4])
|
||||||
|
copy(tcpOut[2:4], tcpIn[0:2])
|
||||||
|
binary.BigEndian.PutUint32(tcpOut[4:], seq)
|
||||||
|
binary.BigEndian.PutUint32(tcpOut[8:], ackSeq)
|
||||||
|
tcpOut[12] = (tcpLen >> 2) << 4 // data offset, reserved, NS
|
||||||
|
tcpOut[13] = outFlags // CWR, ECE, URG, ACK, PSH, RST, SYN, FIN
|
||||||
|
tcpOut[14] = 0 // window size
|
||||||
|
tcpOut[15] = 0 // .
|
||||||
|
tcpOut[16] = 0 // checksum
|
||||||
|
tcpOut[17] = 0 // .
|
||||||
|
tcpOut[18] = 0 // URG Pointer
|
||||||
|
tcpOut[19] = 0 // .
|
||||||
|
|
||||||
|
// Calculate checksum with IPv6 pseudo-header
|
||||||
|
csum := ipv6PseudoheaderChecksum(ipHdr[8:24], ipHdr[24:40], 6, tcpLen)
|
||||||
|
binary.BigEndian.PutUint16(tcpOut[16:], tcpipChecksum(tcpOut, csum))
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
|
||||||
|
nextHeader = packet[6]
|
||||||
|
offset = ipv6.HeaderLen
|
||||||
|
|
||||||
|
for {
|
||||||
|
switch nextHeader {
|
||||||
|
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
|
||||||
|
if len(packet) < offset+2 {
|
||||||
|
return nextHeader, offset, isFragment
|
||||||
|
}
|
||||||
|
nextHeader = packet[offset]
|
||||||
|
offset += (int(packet[offset+1]) + 1) << 3
|
||||||
|
|
||||||
|
case 44: // Fragment
|
||||||
|
if len(packet) < offset+8 {
|
||||||
|
return nextHeader, offset, isFragment
|
||||||
|
}
|
||||||
|
if packet[offset+2] != 0 || packet[offset+3]&0xf8 != 0 {
|
||||||
|
isFragment = true
|
||||||
|
}
|
||||||
|
nextHeader = packet[offset]
|
||||||
|
offset += 8
|
||||||
|
|
||||||
|
case 51: // AH
|
||||||
|
if len(packet) < offset+2 {
|
||||||
|
return nextHeader, offset, isFragment
|
||||||
|
}
|
||||||
|
nextHeader = packet[offset]
|
||||||
|
offset += (int(packet[offset+1]) + 2) << 2
|
||||||
|
|
||||||
|
default:
|
||||||
|
return nextHeader, offset, isFragment
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
||||||
|
if len(packet) < 1 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
switch packet[0] >> 4 {
|
||||||
|
case 4:
|
||||||
|
return createICMPv4EchoResponse(packet, out)
|
||||||
|
case 6:
|
||||||
|
return createICMPv6EchoResponse(packet, out)
|
||||||
|
default:
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func createICMPv4EchoResponse(packet, out []byte) []byte {
|
||||||
// Return early if this is not a simple ICMP Echo Request
|
// Return early if this is not a simple ICMP Echo Request
|
||||||
//TODO: make constants out of these
|
//TODO: make constants out of these
|
||||||
if !(len(packet) >= 28 && len(packet) <= 9001 && packet[0] == 0x45 && packet[9] == 0x01 && packet[20] == 0x08) {
|
if !(len(packet) >= 28 && len(packet) <= 9001 && packet[0] == 0x45 && packet[9] == 0x01 && packet[20] == 0x08) {
|
||||||
@@ -199,6 +418,43 @@ func CreateICMPEchoResponse(packet, out []byte) []byte {
|
|||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func createICMPv6EchoResponse(packet, out []byte) []byte {
|
||||||
|
// IPv6 header (40 bytes) + ICMPv6 header (8 bytes minimum)
|
||||||
|
if len(packet) < ipv6.HeaderLen+8 || len(packet) > 9001 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Next Header must be ICMPv6 (58)
|
||||||
|
if packet[6] != 58 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ICMPv6 type must be Echo Request (128)
|
||||||
|
if packet[ipv6.HeaderLen] != 128 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
out = out[:len(packet)]
|
||||||
|
copy(out, packet)
|
||||||
|
|
||||||
|
// Swap src/dst addresses (bytes 8-23 and 24-39)
|
||||||
|
copy(out[8:24], packet[24:40])
|
||||||
|
copy(out[24:40], packet[8:24])
|
||||||
|
|
||||||
|
// Change ICMPv6 type to Echo Reply (129)
|
||||||
|
icmp := out[ipv6.HeaderLen:]
|
||||||
|
icmp[0] = 129
|
||||||
|
icmp[2] = 0
|
||||||
|
icmp[3] = 0
|
||||||
|
|
||||||
|
// ICMPv6 checksum uses a pseudo-header with src, dst, length, and next header
|
||||||
|
payloadLen := uint32(len(icmp))
|
||||||
|
csum := ipv6PseudoheaderChecksum(out[8:24], out[24:40], 58, payloadLen)
|
||||||
|
binary.BigEndian.PutUint16(icmp[2:], tcpipChecksum(icmp, csum))
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
// calculates the TCP/IP checksum defined in rfc1071. The passed-in
|
// calculates the TCP/IP checksum defined in rfc1071. The passed-in
|
||||||
// csum is any initial checksum data that's already been computed.
|
// csum is any initial checksum data that's already been computed.
|
||||||
//
|
//
|
||||||
@@ -236,3 +492,18 @@ func ipv4PseudoheaderChecksum(src, dst []byte, proto, length uint32) (csum uint3
|
|||||||
csum += length >> 16
|
csum += length >> 16
|
||||||
return csum
|
return csum
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// based on:
|
||||||
|
// - https://github.com/google/gopacket/blob/v1.1.19/layers/tcpip.go#L37-L48
|
||||||
|
func ipv6PseudoheaderChecksum(src, dst []byte, proto, length uint32) (csum uint32) {
|
||||||
|
for i := 0; i < 16; i += 2 {
|
||||||
|
csum += uint32(src[i]) << 8
|
||||||
|
csum += uint32(src[i+1])
|
||||||
|
csum += uint32(dst[i]) << 8
|
||||||
|
csum += uint32(dst[i+1])
|
||||||
|
}
|
||||||
|
csum += proto
|
||||||
|
csum += length & 0xffff
|
||||||
|
csum += length >> 16
|
||||||
|
return csum
|
||||||
|
}
|
||||||
|
|||||||
+445
-1
@@ -1,11 +1,14 @@
|
|||||||
package iputil
|
package iputil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/binary"
|
||||||
"net"
|
"net"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
|
"golang.org/x/net/ipv6"
|
||||||
)
|
)
|
||||||
|
|
||||||
func Test_CreateRejectPacket(t *testing.T) {
|
func Test_CreateRejectPacket(t *testing.T) {
|
||||||
@@ -43,7 +46,7 @@ func Test_CreateRejectPacket(t *testing.T) {
|
|||||||
}
|
}
|
||||||
b = append(b, []byte{0, 3, 0, 4, 0, 0, 0, 0}...)
|
b = append(b, []byte{0, 3, 0, 4, 0, 0, 0, 0}...)
|
||||||
|
|
||||||
expectedLen = MaxRejectPacketSize
|
expectedLen = maxIPv4RejectPacketSize
|
||||||
out = make([]byte, MaxRejectPacketSize)
|
out = make([]byte, MaxRejectPacketSize)
|
||||||
rejectPacket = CreateRejectPacket(b, out)
|
rejectPacket = CreateRejectPacket(b, out)
|
||||||
assert.NotNil(t, rejectPacket)
|
assert.NotNil(t, rejectPacket)
|
||||||
@@ -71,3 +74,444 @@ func Test_CreateRejectPacket(t *testing.T) {
|
|||||||
assert.NotNil(t, rejectPacket)
|
assert.NotNil(t, rejectPacket)
|
||||||
assert.Len(t, rejectPacket, expectedLen)
|
assert.Len(t, rejectPacket, expectedLen)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacket_NoFragment(t *testing.T) {
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
|
||||||
|
// IPv4: non-zero fragment offset should not generate reject packet
|
||||||
|
h := ipv4.Header{
|
||||||
|
Len: 20,
|
||||||
|
Src: net.IPv4(10, 0, 0, 1),
|
||||||
|
Dst: net.IPv4(10, 0, 0, 2),
|
||||||
|
Protocol: 17, // UDP
|
||||||
|
}
|
||||||
|
b, err := h.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("h.Marshal: %v", err)
|
||||||
|
}
|
||||||
|
b = append(b, make([]byte, 8)...)
|
||||||
|
// Set fragment offset to non-zero (byte 6-7, offset in 8-byte units)
|
||||||
|
b[6] = 0x00
|
||||||
|
b[7] = 0x01
|
||||||
|
assert.Nil(t, CreateRejectPacket(b, out))
|
||||||
|
|
||||||
|
// MF flag with zero offset (first fragment) should still generate reject
|
||||||
|
b[6] = 0x20 // MF flag set
|
||||||
|
b[7] = 0x00
|
||||||
|
assert.NotNil(t, CreateRejectPacket(b, out))
|
||||||
|
|
||||||
|
// Non-fragment should still generate reject packet
|
||||||
|
b[6] = 0x00
|
||||||
|
b[7] = 0x00
|
||||||
|
assert.NotNil(t, CreateRejectPacket(b, out))
|
||||||
|
|
||||||
|
// DF flag only (not a fragment) should still generate reject packet
|
||||||
|
b[6] = 0x40
|
||||||
|
b[7] = 0x00
|
||||||
|
assert.NotNil(t, CreateRejectPacket(b, out))
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_NoFragment(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
|
||||||
|
// IPv6 with Fragment header and non-zero offset should not generate reject
|
||||||
|
fragHeader := []byte{
|
||||||
|
17, // next header: UDP
|
||||||
|
0, // reserved
|
||||||
|
0, 9, // fragment offset=1 (shifted left 3), M=1
|
||||||
|
0, 0, 0, 1, // identification
|
||||||
|
}
|
||||||
|
udpPayload := make([]byte, 8)
|
||||||
|
payload := append(fragHeader, udpPayload...)
|
||||||
|
packet := makeIPv6Packet(src, dst, 44, payload) // next header 44 = Fragment
|
||||||
|
assert.Nil(t, CreateRejectPacket(packet, out))
|
||||||
|
|
||||||
|
// Fragment header with zero offset (first fragment) should still generate reject
|
||||||
|
fragHeader[2] = 0
|
||||||
|
fragHeader[3] = 1 // offset=0, M=1
|
||||||
|
payload = append(fragHeader, udpPayload...)
|
||||||
|
packet = makeIPv6Packet(src, dst, 44, payload)
|
||||||
|
assert.NotNil(t, CreateRejectPacket(packet, out))
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacket_NoICMPError(t *testing.T) {
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
|
||||||
|
// ICMP error types should not generate reject packets
|
||||||
|
icmpErrorTypes := []byte{3, 4, 5, 11, 12}
|
||||||
|
for _, icmpType := range icmpErrorTypes {
|
||||||
|
h := ipv4.Header{
|
||||||
|
Len: 20,
|
||||||
|
Src: net.IPv4(10, 0, 0, 1),
|
||||||
|
Dst: net.IPv4(10, 0, 0, 2),
|
||||||
|
Protocol: 1, // ICMP
|
||||||
|
}
|
||||||
|
|
||||||
|
b, err := h.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("h.Marshal: %v", err)
|
||||||
|
}
|
||||||
|
b = append(b, icmpType, 0, 0, 0, 0, 0, 0, 0)
|
||||||
|
|
||||||
|
rejectPacket := CreateRejectPacket(b, out)
|
||||||
|
assert.Nil(t, rejectPacket, "ICMP type %d should not generate a reject packet", icmpType)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ICMP non-error types should still generate reject packets
|
||||||
|
icmpNonErrorTypes := []byte{0, 8, 13, 14}
|
||||||
|
for _, icmpType := range icmpNonErrorTypes {
|
||||||
|
h := ipv4.Header{
|
||||||
|
Len: 20,
|
||||||
|
Src: net.IPv4(10, 0, 0, 1),
|
||||||
|
Dst: net.IPv4(10, 0, 0, 2),
|
||||||
|
Protocol: 1, // ICMP
|
||||||
|
}
|
||||||
|
|
||||||
|
b, err := h.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("h.Marshal: %v", err)
|
||||||
|
}
|
||||||
|
b = append(b, icmpType, 0, 0, 0, 0, 0, 0, 0)
|
||||||
|
|
||||||
|
rejectPacket := CreateRejectPacket(b, out)
|
||||||
|
assert.NotNil(t, rejectPacket, "ICMP type %d should generate a reject packet", icmpType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test_CreateRejectPacket_RespectsCap ensures it is impossible for
|
||||||
|
// an oversized ICMPv6 reject to overwrite the neighbor segment's bytes.
|
||||||
|
func Test_CreateRejectPacket_RespectsCap(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
|
||||||
|
// Inner IPv6 UDP packet. An ICMPv6 reject copies the whole inner packet
|
||||||
|
// plus a 48-byte header (40 IPv6 + 8 ICMPv6), so it needs 48 more bytes
|
||||||
|
// than the inner packet length.
|
||||||
|
inner := makeIPv6Packet(src, dst, 17, make([]byte, 20))
|
||||||
|
|
||||||
|
// The ciphertext scratch reused as the reject buffer is the received
|
||||||
|
// datagram: 16-byte Nebula header + inner + 16-byte AEAD tag. That is only
|
||||||
|
// 32 bytes of slack, so a full ICMPv6 reject overruns it by 16 bytes.
|
||||||
|
const nebulaOverhead = 32
|
||||||
|
segLen := len(inner) + nebulaOverhead
|
||||||
|
|
||||||
|
// Shared backing row laid out as [segment][neighbor's 16-byte Nebula header].
|
||||||
|
const neighborHdr = 16
|
||||||
|
sentinel := bytes.Repeat([]byte{0xAB}, neighborHdr)
|
||||||
|
|
||||||
|
// Uncapped: the slice's capacity reaches into the neighbor, reproducing
|
||||||
|
// the overrun that silently drops the neighbor packet.
|
||||||
|
backing := make([]byte, segLen+neighborHdr)
|
||||||
|
copy(backing[segLen:], sentinel)
|
||||||
|
reject := CreateRejectPacket(inner, backing[:segLen])
|
||||||
|
assert.NotNil(t, reject, "uncapped buffer reaches into the neighbor, so the reject is built")
|
||||||
|
assert.NotEqual(t, sentinel, backing[segLen:segLen+neighborHdr],
|
||||||
|
"without the cap the oversized reject overruns into the neighbor segment")
|
||||||
|
|
||||||
|
// Capped (the fix): cap==len, so the builder cannot exceed the segment. The
|
||||||
|
// reject does not fit, so it is refused rather than corrupting the neighbor.
|
||||||
|
backing = make([]byte, segLen+neighborHdr)
|
||||||
|
copy(backing[segLen:], sentinel)
|
||||||
|
reject = CreateRejectPacket(inner, backing[:segLen:segLen])
|
||||||
|
assert.Nil(t, reject, "capped segment is 16 bytes too small for a full ICMPv6 reject, so it is refused")
|
||||||
|
assert.Equal(t, sentinel, backing[segLen:segLen+neighborHdr],
|
||||||
|
"capped segment must leave the neighbor untouched")
|
||||||
|
}
|
||||||
|
|
||||||
|
func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte {
|
||||||
|
b := make([]byte, ipv6.HeaderLen+len(payload))
|
||||||
|
b[0] = ipv6.Version << 4
|
||||||
|
binary.BigEndian.PutUint16(b[4:], uint16(len(payload)))
|
||||||
|
b[6] = nextHeader
|
||||||
|
b[7] = 64
|
||||||
|
copy(b[8:24], src.To16())
|
||||||
|
copy(b[24:40], dst.To16())
|
||||||
|
copy(b[ipv6.HeaderLen:], payload)
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_ICMP(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
|
||||||
|
// Small UDP packet: entire original included in body
|
||||||
|
udpPayload := make([]byte, 20)
|
||||||
|
udpPayload[0] = 0x00 // src port high
|
||||||
|
udpPayload[1] = 0x50 // src port low (80)
|
||||||
|
udpPayload[2] = 0x01 // dst port high
|
||||||
|
udpPayload[3] = 0xBB // dst port low (443)
|
||||||
|
packet := makeIPv6Packet(src, dst, 17, udpPayload)
|
||||||
|
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.NotNil(t, rejectPacket)
|
||||||
|
|
||||||
|
// Small packet fits entirely: 40 (ipv6 hdr) + 8 (icmpv6 hdr) + 60 (original)
|
||||||
|
expectedLen := ipv6.HeaderLen + 8 + len(packet)
|
||||||
|
assert.Len(t, rejectPacket, expectedLen)
|
||||||
|
|
||||||
|
// Verify version
|
||||||
|
assert.Equal(t, byte(ipv6.Version<<4), rejectPacket[0]&0xf0)
|
||||||
|
// Verify next header is ICMPv6 (58)
|
||||||
|
assert.Equal(t, byte(58), rejectPacket[6])
|
||||||
|
// Verify src/dst are swapped
|
||||||
|
assert.Equal(t, dst.To16(), net.IP(rejectPacket[8:24]))
|
||||||
|
assert.Equal(t, src.To16(), net.IP(rejectPacket[24:40]))
|
||||||
|
// Verify ICMPv6 type=1 (Dest Unreachable), code=1 (Administratively prohibited)
|
||||||
|
assert.Equal(t, byte(1), rejectPacket[ipv6.HeaderLen])
|
||||||
|
assert.Equal(t, byte(1), rejectPacket[ipv6.HeaderLen+1])
|
||||||
|
// Verify entire original packet is included in body
|
||||||
|
assert.Equal(t, packet, rejectPacket[ipv6.HeaderLen+8:])
|
||||||
|
|
||||||
|
// Large packet: body is truncated to 1000 bytes
|
||||||
|
largePkt := makeIPv6Packet(src, dst, 17, make([]byte, 1200))
|
||||||
|
rejectPacket = CreateRejectPacket(largePkt, out)
|
||||||
|
assert.NotNil(t, rejectPacket)
|
||||||
|
assert.Len(t, rejectPacket, ipv6.HeaderLen+8+1000)
|
||||||
|
assert.Equal(t, largePkt[:1000], rejectPacket[ipv6.HeaderLen+8:])
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_TCP(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
|
||||||
|
// TCP SYN packet (next header 6)
|
||||||
|
tcpPayload := make([]byte, 20)
|
||||||
|
tcpPayload[0] = 0x00 // src port high
|
||||||
|
tcpPayload[1] = 0x50 // src port low (80)
|
||||||
|
tcpPayload[2] = 0x01 // dst port high
|
||||||
|
tcpPayload[3] = 0xBB // dst port low (443)
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[4:], 1000) // seq
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[8:], 0) // ack seq
|
||||||
|
tcpPayload[12] = (20 >> 2) << 4 // data offset
|
||||||
|
tcpPayload[13] = 0b00000010 // SYN flag
|
||||||
|
|
||||||
|
packet := makeIPv6Packet(src, dst, 6, tcpPayload)
|
||||||
|
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.NotNil(t, rejectPacket)
|
||||||
|
|
||||||
|
// Expected: 40 (ipv6 hdr) + 20 (tcp RST)
|
||||||
|
expectedLen := ipv6.HeaderLen + 20
|
||||||
|
assert.Len(t, rejectPacket, expectedLen)
|
||||||
|
|
||||||
|
// Verify version
|
||||||
|
assert.Equal(t, byte(ipv6.Version<<4), rejectPacket[0]&0xf0)
|
||||||
|
// Verify next header is TCP (6)
|
||||||
|
assert.Equal(t, byte(6), rejectPacket[6])
|
||||||
|
// Verify src/dst are swapped
|
||||||
|
assert.Equal(t, dst.To16(), net.IP(rejectPacket[8:24]))
|
||||||
|
assert.Equal(t, src.To16(), net.IP(rejectPacket[24:40]))
|
||||||
|
// Verify ports are swapped
|
||||||
|
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
||||||
|
assert.Equal(t, uint16(443), binary.BigEndian.Uint16(tcpOut[0:2]))
|
||||||
|
assert.Equal(t, uint16(80), binary.BigEndian.Uint16(tcpOut[2:4]))
|
||||||
|
// RST+ACK flags (since input was SYN without ACK)
|
||||||
|
assert.Equal(t, byte(0b00010100), tcpOut[13])
|
||||||
|
// ack_seq = original seq (1000) + SYN (1) + FIN (0) + segment data (0)
|
||||||
|
assert.Equal(t, uint32(1001), binary.BigEndian.Uint32(tcpOut[8:]))
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_TCPWithACK(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
|
||||||
|
// TCP packet with ACK set
|
||||||
|
tcpPayload := make([]byte, 20)
|
||||||
|
tcpPayload[0] = 0x00
|
||||||
|
tcpPayload[1] = 0x50
|
||||||
|
tcpPayload[2] = 0x01
|
||||||
|
tcpPayload[3] = 0xBB
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[4:], 1000) // seq
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[8:], 2000) // ack seq
|
||||||
|
tcpPayload[12] = (20 >> 2) << 4 // data offset
|
||||||
|
tcpPayload[13] = 0b00010000 // ACK flag
|
||||||
|
|
||||||
|
packet := makeIPv6Packet(src, dst, 6, tcpPayload)
|
||||||
|
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.NotNil(t, rejectPacket)
|
||||||
|
|
||||||
|
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
||||||
|
// RST only (no ACK) since input had ACK
|
||||||
|
assert.Equal(t, byte(0b00000100), tcpOut[13])
|
||||||
|
// seq = original ack_seq
|
||||||
|
assert.Equal(t, uint32(2000), binary.BigEndian.Uint32(tcpOut[4:]))
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_NoICMPError(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
|
||||||
|
// ICMPv6 error types (1-4) should not generate reject packets
|
||||||
|
for icmpType := byte(1); icmpType <= 4; icmpType++ {
|
||||||
|
payload := make([]byte, 8)
|
||||||
|
payload[0] = icmpType
|
||||||
|
packet := makeIPv6Packet(src, dst, 58, payload)
|
||||||
|
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.Nil(t, rejectPacket, "ICMPv6 type %d should not generate a reject packet", icmpType)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ICMPv6 non-error types should still generate reject packets
|
||||||
|
nonErrorTypes := []byte{128, 129, 133, 134}
|
||||||
|
for _, icmpType := range nonErrorTypes {
|
||||||
|
payload := make([]byte, 8)
|
||||||
|
payload[0] = icmpType
|
||||||
|
packet := makeIPv6Packet(src, dst, 58, payload)
|
||||||
|
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.NotNil(t, rejectPacket, "ICMPv6 type %d should generate a reject packet", icmpType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_TooShort(t *testing.T) {
|
||||||
|
// Packet too short to be valid IPv6
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
assert.Nil(t, CreateRejectPacket([]byte{0x60}, out))
|
||||||
|
assert.Nil(t, CreateRejectPacket(make([]byte, 39), out))
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_ExtensionHeaders(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
|
||||||
|
// IPv6 + Hop-by-Hop extension header + TCP
|
||||||
|
hopByHop := []byte{
|
||||||
|
6, // next header: TCP
|
||||||
|
0, // length (8 bytes total)
|
||||||
|
0, 0, // padding
|
||||||
|
0, 0, 0, 0,
|
||||||
|
}
|
||||||
|
tcpPayload := make([]byte, 20)
|
||||||
|
tcpPayload[0] = 0x00
|
||||||
|
tcpPayload[1] = 0x50
|
||||||
|
tcpPayload[2] = 0x01
|
||||||
|
tcpPayload[3] = 0xBB
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[4:], 1000)
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[8:], 2000)
|
||||||
|
tcpPayload[12] = (20 >> 2) << 4
|
||||||
|
tcpPayload[13] = 0b00010000 // ACK
|
||||||
|
|
||||||
|
payload := append(hopByHop, tcpPayload...)
|
||||||
|
packet := makeIPv6Packet(src, dst, 0, payload) // next header 0 = Hop-by-Hop
|
||||||
|
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.NotNil(t, rejectPacket)
|
||||||
|
|
||||||
|
// Should produce TCP RST
|
||||||
|
expectedLen := ipv6.HeaderLen + 20
|
||||||
|
assert.Len(t, rejectPacket, expectedLen)
|
||||||
|
assert.Equal(t, byte(6), rejectPacket[6]) // next header is TCP
|
||||||
|
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
||||||
|
assert.Equal(t, byte(0b00000100), tcpOut[13]) // RST only
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateICMPEchoResponse_IPv4(t *testing.T) {
|
||||||
|
// Build a simple IPv4 ICMP Echo Request
|
||||||
|
packet := make([]byte, 28)
|
||||||
|
packet[0] = 0x45 // version 4, IHL 5
|
||||||
|
binary.BigEndian.PutUint16(packet[2:], uint16(28)) // total length
|
||||||
|
packet[8] = 64 // TTL
|
||||||
|
packet[9] = 1 // protocol ICMP
|
||||||
|
copy(packet[12:16], net.IPv4(10, 0, 0, 1).To4()) // src
|
||||||
|
copy(packet[16:20], net.IPv4(10, 0, 0, 2).To4()) // dst
|
||||||
|
packet[20] = 8 // ICMP Echo Request
|
||||||
|
|
||||||
|
out := make([]byte, len(packet))
|
||||||
|
result := CreateICMPEchoResponse(packet, out)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.Equal(t, byte(0x45), result[0])
|
||||||
|
// src/dst swapped
|
||||||
|
assert.Equal(t, net.IPv4(10, 0, 0, 2).To4(), net.IP(result[12:16]))
|
||||||
|
assert.Equal(t, net.IPv4(10, 0, 0, 1).To4(), net.IP(result[16:20]))
|
||||||
|
// ICMP Echo Reply
|
||||||
|
assert.Equal(t, byte(0), result[20])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateICMPEchoResponse_IPv6(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1").To16()
|
||||||
|
dst := net.ParseIP("fd00::2").To16()
|
||||||
|
|
||||||
|
// Build an IPv6 ICMPv6 Echo Request packet
|
||||||
|
// IPv6 header (40 bytes) + ICMPv6 (8 bytes)
|
||||||
|
packet := make([]byte, 48)
|
||||||
|
packet[0] = 0x60 // version 6
|
||||||
|
payloadLen := uint16(8) // ICMPv6 header only
|
||||||
|
binary.BigEndian.PutUint16(packet[4:], payloadLen)
|
||||||
|
packet[6] = 58 // Next Header: ICMPv6
|
||||||
|
packet[7] = 64 // Hop Limit
|
||||||
|
copy(packet[8:24], src) // src address
|
||||||
|
copy(packet[24:40], dst) // dst address
|
||||||
|
|
||||||
|
// ICMPv6 Echo Request
|
||||||
|
icmp := packet[40:]
|
||||||
|
icmp[0] = 128 // type: Echo Request
|
||||||
|
icmp[1] = 0 // code
|
||||||
|
binary.BigEndian.PutUint16(icmp[4:], 1) // identifier
|
||||||
|
binary.BigEndian.PutUint16(icmp[6:], 1) // sequence number
|
||||||
|
|
||||||
|
// Compute correct checksum for the request
|
||||||
|
csum := ipv6PseudoheaderChecksum(src, dst, 58, uint32(payloadLen))
|
||||||
|
binary.BigEndian.PutUint16(icmp[2:], tcpipChecksum(icmp, csum))
|
||||||
|
|
||||||
|
out := make([]byte, len(packet))
|
||||||
|
result := CreateICMPEchoResponse(packet, out)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
|
||||||
|
// Version should still be 6
|
||||||
|
assert.Equal(t, byte(6), result[0]>>4)
|
||||||
|
// src/dst swapped
|
||||||
|
assert.Equal(t, dst, net.IP(result[8:24]))
|
||||||
|
assert.Equal(t, src, net.IP(result[24:40]))
|
||||||
|
// ICMPv6 Echo Reply type
|
||||||
|
assert.Equal(t, byte(129), result[40])
|
||||||
|
|
||||||
|
// Verify checksum is valid (tcpipChecksum returns 0 when data+checksum is correct)
|
||||||
|
respIcmp := result[40:]
|
||||||
|
verifyCsum := ipv6PseudoheaderChecksum(result[8:24], result[24:40], 58, uint32(payloadLen))
|
||||||
|
assert.Equal(t, uint16(0), tcpipChecksum(respIcmp, verifyCsum))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateICMPEchoResponse_IPv6_NotEchoRequest(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1").To16()
|
||||||
|
dst := net.ParseIP("fd00::2").To16()
|
||||||
|
|
||||||
|
packet := make([]byte, 48)
|
||||||
|
packet[0] = 0x60
|
||||||
|
binary.BigEndian.PutUint16(packet[4:], 8)
|
||||||
|
packet[6] = 58
|
||||||
|
packet[7] = 64
|
||||||
|
copy(packet[8:24], src)
|
||||||
|
copy(packet[24:40], dst)
|
||||||
|
|
||||||
|
// ICMPv6 type 1 (Destination Unreachable) - not Echo Request
|
||||||
|
packet[40] = 1
|
||||||
|
|
||||||
|
out := make([]byte, len(packet))
|
||||||
|
result := CreateICMPEchoResponse(packet, out)
|
||||||
|
assert.Nil(t, result)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1").To16()
|
||||||
|
dst := net.ParseIP("fd00::2").To16()
|
||||||
|
|
||||||
|
packet := make([]byte, 48)
|
||||||
|
packet[0] = 0x60
|
||||||
|
binary.BigEndian.PutUint16(packet[4:], 8)
|
||||||
|
packet[6] = 6 // TCP, not ICMPv6
|
||||||
|
packet[7] = 64
|
||||||
|
copy(packet[8:24], src)
|
||||||
|
copy(packet[24:40], dst)
|
||||||
|
|
||||||
|
out := make([]byte, len(packet))
|
||||||
|
result := CreateICMPEchoResponse(packet, out)
|
||||||
|
assert.Nil(t, result)
|
||||||
|
}
|
||||||
|
|||||||
+37
-52
@@ -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)
|
||||||
@@ -279,16 +272,18 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
|
|||||||
//NOTE: many things will get much simpler when we combine static_host_map and lighthouse.hosts in config
|
//NOTE: many things will get much simpler when we combine static_host_map and lighthouse.hosts in config
|
||||||
if initial || c.HasChanged("static_host_map") || c.HasChanged("static_map.cadence") || c.HasChanged("static_map.network") || c.HasChanged("static_map.lookup_timeout") {
|
if initial || c.HasChanged("static_host_map") || c.HasChanged("static_map.cadence") || c.HasChanged("static_map.network") || c.HasChanged("static_map.lookup_timeout") {
|
||||||
// Clean up. Entries still in the static_host_map will be re-built.
|
// Clean up. Entries still in the static_host_map will be re-built.
|
||||||
// Entries no longer present must have their (possible) background DNS goroutines stopped.
|
ourselves := lh.myVpnNetworks[0].Addr()
|
||||||
if existingStaticList := lh.staticList.Load(); existingStaticList != nil {
|
oldStaticList := lh.staticList.Load()
|
||||||
|
if oldStaticList != nil {
|
||||||
lh.RLock()
|
lh.RLock()
|
||||||
for staticVpnAddr := range *existingStaticList {
|
for staticVpnAddr := range *oldStaticList {
|
||||||
if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil {
|
if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil {
|
||||||
am.hr.Cancel()
|
am.ResetForOwner(ourselves)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
lh.RUnlock()
|
lh.RUnlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build a new list based on current config.
|
// Build a new list based on current config.
|
||||||
staticList := make(map[netip.Addr]struct{})
|
staticList := make(map[netip.Addr]struct{})
|
||||||
err := lh.loadStaticMap(c, staticList)
|
err := lh.loadStaticMap(c, staticList)
|
||||||
@@ -296,6 +291,21 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// For entries removed from static_host_map, stop the DNS goroutine and drop the cached addrs.
|
||||||
|
// All addrs must come from the lighthouses now that it's no longer a static host.
|
||||||
|
if oldStaticList != nil {
|
||||||
|
lh.RLock()
|
||||||
|
for staticVpnAddr := range *oldStaticList {
|
||||||
|
if _, stillStatic := staticList[staticVpnAddr]; stillStatic {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil {
|
||||||
|
am.ClearHostnameResults()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
lh.RUnlock()
|
||||||
|
}
|
||||||
|
|
||||||
lh.staticList.Store(&staticList)
|
lh.staticList.Store(&staticList)
|
||||||
if !initial {
|
if !initial {
|
||||||
if c.HasChanged("static_host_map") {
|
if c.HasChanged("static_host_map") {
|
||||||
@@ -1406,58 +1416,31 @@ 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 {
|
||||||
|
if a == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
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 {
|
||||||
|
if a == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
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 {
|
||||||
@@ -1477,7 +1460,7 @@ func protoV6AddrPortToNetAddrPort(ap *V6AddrPort) netip.AddrPort {
|
|||||||
b := [16]byte{}
|
b := [16]byte{}
|
||||||
binary.BigEndian.PutUint64(b[:8], ap.Hi)
|
binary.BigEndian.PutUint64(b[:8], ap.Hi)
|
||||||
binary.BigEndian.PutUint64(b[8:], ap.Lo)
|
binary.BigEndian.PutUint64(b[8:], ap.Lo)
|
||||||
return netip.AddrPortFrom(netip.AddrFrom16(b), uint16(ap.Port))
|
return netip.AddrPortFrom(netip.AddrFrom16(b).Unmap(), uint16(ap.Port))
|
||||||
}
|
}
|
||||||
|
|
||||||
func netAddrToProtoAddr(addr netip.Addr) *Addr {
|
func netAddrToProtoAddr(addr netip.Addr) *Addr {
|
||||||
@@ -1517,7 +1500,9 @@ func (d *NebulaMetaDetails) GetRelays() []netip.Addr {
|
|||||||
|
|
||||||
if len(d.RelayVpnAddrs) > 0 {
|
if len(d.RelayVpnAddrs) > 0 {
|
||||||
for _, r := range d.RelayVpnAddrs {
|
for _, r := range d.RelayVpnAddrs {
|
||||||
relays = append(relays, protoAddrToNetAddr(r))
|
if r != nil {
|
||||||
|
relays = append(relays, protoAddrToNetAddr(r))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return relays
|
return relays
|
||||||
|
|||||||
@@ -303,6 +303,132 @@ func TestLighthouse_reload(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestLighthouse_reloadStaticHostMap verifies that reloading static_host_map applies the new
|
||||||
|
// config rather than appending to it. See issue #718.
|
||||||
|
func TestLighthouse_reloadStaticHostMap(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
c := config.NewC(l)
|
||||||
|
c.Settings["lighthouse"] = map[string]any{"am_lighthouse": true}
|
||||||
|
c.Settings["listen"] = map[string]any{"port": 4242}
|
||||||
|
c.Settings["static_host_map"] = map[string]any{
|
||||||
|
"10.128.0.2": []any{"1.1.1.1:4242"},
|
||||||
|
}
|
||||||
|
|
||||||
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||||
|
nt := new(bart.Lite)
|
||||||
|
nt.Insert(myVpnNet)
|
||||||
|
cs := &CertState{
|
||||||
|
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||||
|
myVpnNetworksTable: nt,
|
||||||
|
}
|
||||||
|
|
||||||
|
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
staticHost := netip.MustParseAddr("10.128.0.2")
|
||||||
|
otherHost := netip.MustParseAddr("10.128.0.3")
|
||||||
|
|
||||||
|
// Capture the RemoteList pointer up front; an in-flight handshake would hold the same one
|
||||||
|
// on hostinfo.remotes, so it must reflect every reload below.
|
||||||
|
pinned := lh.Query(staticHost)
|
||||||
|
require.NotNil(t, pinned)
|
||||||
|
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, pinned.CopyAddrs([]netip.Prefix{}))
|
||||||
|
|
||||||
|
// Replace the remote address. The new address should be the only entry.
|
||||||
|
nc := map[string]any{
|
||||||
|
"static_host_map": map[string]any{
|
||||||
|
"10.128.0.2": []any{"2.2.2.2:4242"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
rc, err := yaml.Marshal(nc)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, c.ReloadConfigString(string(rc)))
|
||||||
|
|
||||||
|
rl := lh.Query(staticHost)
|
||||||
|
require.NotNil(t, rl)
|
||||||
|
assert.Same(t, pinned, rl, "RemoteList pointer must stay stable so in-flight handshakes pick up the change")
|
||||||
|
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("2.2.2.2:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
||||||
|
|
||||||
|
// Reload back to the original IP. Mirrors the round-trip in issue #718 step 6-8 where
|
||||||
|
// the buggy reload produced [1.1.1.1, 2.2.2.2, 1.1.1.1] instead of [1.1.1.1].
|
||||||
|
nc = map[string]any{
|
||||||
|
"static_host_map": map[string]any{
|
||||||
|
"10.128.0.2": []any{"1.1.1.1:4242"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
rc, err = yaml.Marshal(nc)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, c.ReloadConfigString(string(rc)))
|
||||||
|
|
||||||
|
rl = lh.Query(staticHost)
|
||||||
|
require.NotNil(t, rl)
|
||||||
|
assert.Same(t, pinned, rl)
|
||||||
|
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
||||||
|
|
||||||
|
// Reload with the same config. An unchanged entry must not duplicate.
|
||||||
|
require.NoError(t, c.ReloadConfigString(string(rc)))
|
||||||
|
|
||||||
|
rl = lh.Query(staticHost)
|
||||||
|
require.NotNil(t, rl)
|
||||||
|
assert.Same(t, pinned, rl)
|
||||||
|
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
||||||
|
|
||||||
|
// Switch back to 2.2.2.2 so the rest of the test continues against a known address.
|
||||||
|
nc = map[string]any{
|
||||||
|
"static_host_map": map[string]any{
|
||||||
|
"10.128.0.2": []any{"2.2.2.2:4242"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
rc, err = yaml.Marshal(nc)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, c.ReloadConfigString(string(rc)))
|
||||||
|
|
||||||
|
// Add a second host alongside the first. Both should be present, neither duplicated.
|
||||||
|
nc = map[string]any{
|
||||||
|
"static_host_map": map[string]any{
|
||||||
|
"10.128.0.2": []any{"2.2.2.2:4242"},
|
||||||
|
"10.128.0.3": []any{"3.3.3.3:4242"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
rc, err = yaml.Marshal(nc)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, c.ReloadConfigString(string(rc)))
|
||||||
|
|
||||||
|
rl = lh.Query(staticHost)
|
||||||
|
require.NotNil(t, rl)
|
||||||
|
assert.Same(t, pinned, rl, "adding a sibling entry must not displace the existing RemoteList")
|
||||||
|
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("2.2.2.2:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
||||||
|
|
||||||
|
rl = lh.Query(otherHost)
|
||||||
|
require.NotNil(t, rl)
|
||||||
|
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
||||||
|
|
||||||
|
// Drop the first host entirely. The vpnAddr is no longer marked static, our owner
|
||||||
|
// contribution is cleared, but the addrMap entry stays in place so non-static cache
|
||||||
|
// data (from lighthouse queries) on the same RemoteList isn't lost. In-flight handshakes
|
||||||
|
// that already had the pointer see an empty address list rather than retrying stale ones.
|
||||||
|
nc = map[string]any{
|
||||||
|
"static_host_map": map[string]any{
|
||||||
|
"10.128.0.3": []any{"3.3.3.3:4242"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
rc, err = yaml.Marshal(nc)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, c.ReloadConfigString(string(rc)))
|
||||||
|
|
||||||
|
_, isStatic := lh.GetStaticHostList()[staticHost]
|
||||||
|
assert.False(t, isStatic)
|
||||||
|
|
||||||
|
rl = lh.Query(staticHost)
|
||||||
|
require.NotNil(t, rl)
|
||||||
|
assert.Same(t, pinned, rl)
|
||||||
|
assert.Empty(t, rl.CopyAddrs([]netip.Prefix{}))
|
||||||
|
|
||||||
|
rl = lh.Query(otherHost)
|
||||||
|
require.NotNil(t, rl)
|
||||||
|
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
||||||
|
}
|
||||||
|
|
||||||
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
||||||
req := &NebulaMeta{
|
req := &NebulaMeta{
|
||||||
Type: NebulaMeta_HostQuery,
|
Type: NebulaMeta_HostQuery,
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -33,6 +34,9 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
buildVersion = moduleVersion()
|
buildVersion = moduleVersion()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Debug builds (-tags debug) serve pprof on :6060; a no-op otherwise.
|
||||||
|
startPprofServer(ctx, l)
|
||||||
|
|
||||||
// Print the config if in test, the exit comes later
|
// Print the config if in test, the exit comes later
|
||||||
if configTest {
|
if configTest {
|
||||||
b, err := yaml.Marshal(c.Settings)
|
b, err := yaml.Marshal(c.Settings)
|
||||||
@@ -55,7 +59,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)
|
||||||
}
|
}
|
||||||
@@ -130,6 +134,17 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
udpConns := make([]udp.Conn, routines)
|
udpConns := make([]udp.Conn, routines)
|
||||||
port := c.GetInt("listen.port", 0)
|
port := c.GetInt("listen.port", 0)
|
||||||
|
|
||||||
|
// Callers get no handle to these until the Control is returned, release them on any error.
|
||||||
|
defer func() {
|
||||||
|
if reterr != nil {
|
||||||
|
for _, u := range udpConns {
|
||||||
|
if u != nil {
|
||||||
|
_ = u.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
if !configTest {
|
if !configTest {
|
||||||
rawListenHost := c.GetString("listen.host", "0.0.0.0")
|
rawListenHost := c.GetString("listen.host", "0.0.0.0")
|
||||||
var listenHost netip.Addr
|
var listenHost netip.Addr
|
||||||
@@ -170,7 +185,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 {
|
||||||
@@ -194,11 +209,17 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
handshakeManager := NewHandshakeManager(l, hostMap, lightHouse, udpConns[0], handshakeConfig)
|
handshakeManager := NewHandshakeManager(l, hostMap, lightHouse, udpConns[0], handshakeConfig)
|
||||||
lightHouse.handshakeTrigger = handshakeManager.trigger
|
lightHouse.handshakeTrigger = handshakeManager.trigger
|
||||||
|
|
||||||
ds, err := newDnsServerFromConfig(ctx, l, pki.getCertState(), hostMap, c)
|
ds, err := newDnsServerFromConfig(ctx, l, pki, hostMap, c)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
l.Warn("Failed to start DNS responder", "error", err)
|
l.Warn("Failed to start DNS responder", "error", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pinThreads := c.GetBool("tun.pin_threads", true)
|
||||||
|
cpuAffinity := parseCpuAffinity(c, l, routines)
|
||||||
|
if pinThreads && len(cpuAffinity) == 0 && !configTest {
|
||||||
|
cpuAffinity = defaultCPUAffinityAvoidingIRQs(l, routines)
|
||||||
|
}
|
||||||
|
|
||||||
ifConfig := &InterfaceConfig{
|
ifConfig := &InterfaceConfig{
|
||||||
HostMap: hostMap,
|
HostMap: hostMap,
|
||||||
Inside: tun,
|
Inside: tun,
|
||||||
@@ -220,6 +241,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
||||||
punchy: punchy,
|
punchy: punchy,
|
||||||
ConntrackCacheTimeout: conntrackCacheTimeout,
|
ConntrackCacheTimeout: conntrackCacheTimeout,
|
||||||
|
CpuAffinity: cpuAffinity,
|
||||||
|
PinThreads: pinThreads,
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -237,9 +260,12 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
ifce.reloadDisconnectInvalid(c)
|
ifce.reloadDisconnectInvalid(c)
|
||||||
ifce.reloadSendRecvError(c)
|
ifce.reloadSendRecvError(c)
|
||||||
ifce.reloadAcceptRecvError(c)
|
ifce.reloadAcceptRecvError(c)
|
||||||
|
ifce.reloadEcn(c)
|
||||||
|
|
||||||
handshakeManager.f = ifce
|
handshakeManager.f = ifce
|
||||||
go handshakeManager.Run(ctx)
|
go handshakeManager.Run(ctx)
|
||||||
|
|
||||||
|
punchy.Start(ctx, ifce, hostMap, lightHouse)
|
||||||
}
|
}
|
||||||
|
|
||||||
stats, err := newStatsServerFromConfig(ctx, l, c, buildVersion, configTest)
|
stats, err := newStatsServerFromConfig(ctx, l, c, buildVersion, configTest)
|
||||||
@@ -269,6 +295,121 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// parseCpuAffinity reads `tun.cpu_affinity` from the config — a list of
|
||||||
|
// integer CPU IDs, one per TUN reader goroutine. Empty / unset returns nil
|
||||||
|
// (listenIn falls back to spreading queues across the allowed CPU set).
|
||||||
|
// Length mismatch with `routines` is a warning, not an error: shorter lists
|
||||||
|
// are modulo-cycled across queues, longer lists' tail is ignored. Invalid
|
||||||
|
// entries (non-integer, or a CPU ID we're not allowed to run on) are also a
|
||||||
|
// warning and disable the override entirely so we don't silently pin to the
|
||||||
|
// wrong CPU. Entries are validated against the process's current affinity
|
||||||
|
// mask (util.AllowedCPUs) rather than 0..NumCPU-1: under a cgroup cpuset or
|
||||||
|
// taskset the runnable IDs are frequently not that contiguous range, and
|
||||||
|
// pinning to an unrunnable ID always fails. If the allowed set can't be
|
||||||
|
// determined we fall back to a plain non-negative check.
|
||||||
|
func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
|
||||||
|
raw := c.Get("tun.cpu_affinity")
|
||||||
|
if raw == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
rv, ok := raw.([]any)
|
||||||
|
if !ok {
|
||||||
|
l.Warn("tun.cpu_affinity must be a list of integers; ignoring", "value", raw)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// allowed is the set of CPU IDs we're actually permitted to run on. A nil
|
||||||
|
// slice (unsupported platform or lookup error) means "can't tell", so we
|
||||||
|
// only apply the weaker non-negative check in that case.
|
||||||
|
allowed, err := util.AllowedCPUs()
|
||||||
|
if err != nil {
|
||||||
|
l.Warn("could not determine allowed CPUs; validating tun.cpu_affinity against non-negative only", "error", err)
|
||||||
|
allowed = nil
|
||||||
|
}
|
||||||
|
cpus := make([]int, 0, len(rv))
|
||||||
|
for i, e := range rv {
|
||||||
|
var cpu int
|
||||||
|
switch v := e.(type) {
|
||||||
|
case int:
|
||||||
|
cpu = v
|
||||||
|
case int64:
|
||||||
|
cpu = int(v)
|
||||||
|
case float64:
|
||||||
|
cpu = int(v)
|
||||||
|
default:
|
||||||
|
l.Warn("tun.cpu_affinity entry not an integer; ignoring affinity",
|
||||||
|
"index", i, "value", e)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if cpu < 0 {
|
||||||
|
l.Warn("tun.cpu_affinity entry out of range; ignoring affinity",
|
||||||
|
"index", i, "cpu", cpu)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if len(allowed) > 0 && !slices.Contains(allowed, cpu) {
|
||||||
|
l.Warn("tun.cpu_affinity entry not in allowed CPU set; ignoring affinity",
|
||||||
|
"index", i, "cpu", cpu, "allowed", allowed)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
cpus = append(cpus, cpu)
|
||||||
|
}
|
||||||
|
if len(cpus) != routines {
|
||||||
|
l.Warn("tun.cpu_affinity length doesn't match routines; queues will modulo-cycle through the list",
|
||||||
|
"affinity_len", len(cpus), "routines", routines)
|
||||||
|
}
|
||||||
|
return cpus
|
||||||
|
}
|
||||||
|
|
||||||
|
// defaultCPUAffinityAvoidingIRQs picks the default pin set for the tun
|
||||||
|
// readers when tun.cpu_affinity is unset: allowed CPUs that do NOT service
|
||||||
|
// any physical NIC's interrupts. The stock allowed[i] spread pins the
|
||||||
|
// encrypt threads onto exactly the cores most drivers affine their first RX
|
||||||
|
// queue IRQs to, so whenever a flow's RSS queue fires on a core hosting a
|
||||||
|
// tun reader, NAPI and encrypt fight for the core and per-flow throughput
|
||||||
|
// drops (measured: REV 8.4 vs 10.2 Gbps on the same hardware, 2026-07-14).
|
||||||
|
//
|
||||||
|
// Returns nil — keeping the old allowed[i] fallback in listenIn — when IRQ
|
||||||
|
// info is unavailable or when there aren't enough IRQ-free CPUs to give
|
||||||
|
// every routine its own core: silently doubling readers up on fewer cores
|
||||||
|
// is worse than the occasional IRQ collision. NICs whose vectors blanket
|
||||||
|
// every CPU (e.g. mlx5 defaults to one queue per core) make avoidance
|
||||||
|
// impossible; narrowing the NIC's spread (ethtool -X <dev> equal N, or
|
||||||
|
// /proc/irq/*/smp_affinity) or setting tun.cpu_affinity explicitly makes it
|
||||||
|
// effective.
|
||||||
|
func defaultCPUAffinityAvoidingIRQs(l *slog.Logger, routines int) []int {
|
||||||
|
irq, err := util.NICIRQCPUs()
|
||||||
|
if err != nil || len(irq) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
allowed, err := util.AllowedCPUs()
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
cpus := chooseIRQFreeCPUs(allowed, irq, routines)
|
||||||
|
if cpus == nil {
|
||||||
|
l.Info("not enough CPUs are free of NIC IRQs to give every tun reader its own; using the default spread",
|
||||||
|
"routines", routines, "allowed", len(allowed), "irqCPUs", len(irq))
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
l.Info("pinning tun readers to CPUs clear of NIC IRQs", "cpus", cpus)
|
||||||
|
return cpus
|
||||||
|
}
|
||||||
|
|
||||||
|
// chooseIRQFreeCPUs returns the first `routines` allowed CPUs not present in
|
||||||
|
// irq, or nil if fewer than `routines` qualify.
|
||||||
|
func chooseIRQFreeCPUs(allowed []int, irq map[int]bool, routines int) []int {
|
||||||
|
free := make([]int, 0, routines)
|
||||||
|
for _, cpu := range allowed {
|
||||||
|
if irq[cpu] {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
free = append(free, cpu)
|
||||||
|
if len(free) == routines {
|
||||||
|
return free
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func moduleVersion() string {
|
func moduleVersion() string {
|
||||||
info, ok := debug.ReadBuildInfo()
|
info, ok := debug.ReadBuildInfo()
|
||||||
if !ok {
|
if !ok {
|
||||||
|
|||||||
@@ -0,0 +1,71 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestChooseIRQFreeCPUs(t *testing.T) {
|
||||||
|
irq := map[int]bool{0: true, 1: true, 2: true, 3: true}
|
||||||
|
|
||||||
|
// Plenty of IRQ-free CPUs: take the first `routines` of them in order.
|
||||||
|
assert.Equal(t, []int{4, 5}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 2))
|
||||||
|
|
||||||
|
// Exactly enough.
|
||||||
|
assert.Equal(t, []int{4, 5, 6}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 3))
|
||||||
|
|
||||||
|
// Not enough IRQ-free CPUs: nil, caller keeps the old default rather
|
||||||
|
// than doubling readers up on shared cores.
|
||||||
|
assert.Nil(t, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4}, irq, 2))
|
||||||
|
|
||||||
|
// No IRQ info at all behaves like a plain prefix of allowed.
|
||||||
|
assert.Equal(t, []int{0, 1}, chooseIRQFreeCPUs([]int{0, 1, 2}, map[int]bool{}, 2))
|
||||||
|
|
||||||
|
// Non-contiguous allowed set (cgroup cpuset) with holes.
|
||||||
|
assert.Equal(t, []int{9, 12}, chooseIRQFreeCPUs([]int{1, 3, 9, 12}, map[int]bool{1: true, 3: true}, 2))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseCpuAffinity(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
|
||||||
|
// newConfig returns a config.C with tun.cpu_affinity set to v. A nil v
|
||||||
|
// leaves the key unset.
|
||||||
|
newConfig := func(v any) *config.C {
|
||||||
|
c := config.NewC(l)
|
||||||
|
if v != nil {
|
||||||
|
c.Settings["tun"] = map[string]any{"cpu_affinity": v}
|
||||||
|
}
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// unset -> nil (listenIn falls back to spreading across the allowed set)
|
||||||
|
assert.Nil(t, parseCpuAffinity(newConfig(nil), l, 1))
|
||||||
|
|
||||||
|
// Pick a CPU we're actually allowed to run on so a valid list survives
|
||||||
|
// validation regardless of the host's affinity mask.
|
||||||
|
allowed, _ := util.AllowedCPUs()
|
||||||
|
validCPU := 0
|
||||||
|
if len(allowed) > 0 {
|
||||||
|
validCPU = allowed[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
// valid list -> parsed through unchanged
|
||||||
|
assert.Equal(t, []int{validCPU, validCPU}, parseCpuAffinity(newConfig([]any{validCPU, validCPU}), l, 2))
|
||||||
|
|
||||||
|
// a negative entry is out of range on every platform -> disables the override
|
||||||
|
assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, -1}), l, 2))
|
||||||
|
|
||||||
|
// a non-integer entry -> disables the override
|
||||||
|
assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, "not-a-cpu"}), l, 2))
|
||||||
|
|
||||||
|
// a CPU id outside the allowed set -> disables the override. Only assertable
|
||||||
|
// where we can enumerate the allowed set (e.g. linux); 1<<20 is far beyond
|
||||||
|
// any representable CPU id so it can never be in the mask.
|
||||||
|
if len(allowed) > 0 {
|
||||||
|
assert.Nil(t, parseCpuAffinity(newConfig([]any{1 << 20}), l, 1))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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,166 @@
|
|||||||
|
package noiseutil
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCipherStateAESGCMRoundtrip(t *testing.T) {
|
||||||
|
enc, dec := buildCipherStates(t, CipherAESGCM)
|
||||||
|
roundtrip(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCipherStateChaChaPolyRoundtrip(t *testing.T) {
|
||||||
|
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||||
|
roundtrip(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewCipherStateDispatch(t *testing.T) {
|
||||||
|
encA, _ := buildCipherStates(t, CipherAESGCM)
|
||||||
|
encC, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||||
|
|
||||||
|
assert.IsType(t, &CipherStateAESGCM{}, NewCipherState(encA, CipherAESGCM))
|
||||||
|
assert.IsType(t, &CipherStateChaChaPoly{}, NewCipherState(encC, noise.CipherChaChaPoly))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewCipherStateUnsupportedPanics(t *testing.T) {
|
||||||
|
enc, _ := buildCipherStates(t, CipherAESGCM)
|
||||||
|
assert.Panics(t, func() {
|
||||||
|
NewCipherState(enc, fakeCipher{})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
type fakeCipher struct{}
|
||||||
|
|
||||||
|
func (fakeCipher) Cipher(k [32]byte) noise.Cipher { return nil }
|
||||||
|
func (fakeCipher) CipherName() string { return "Fake" }
|
||||||
|
|
||||||
|
// buildCipherStates runs an in-memory NN handshake with the requested cipher
|
||||||
|
// to produce a pair of post-handshake CipherStates that share keys.
|
||||||
|
func buildCipherStates(t *testing.T, c noise.CipherFunc) (*noise.CipherState, *noise.CipherState) {
|
||||||
|
t.Helper()
|
||||||
|
suite := noise.NewCipherSuite(noise.DH25519, c, noise.HashSHA256)
|
||||||
|
cfg := noise.Config{CipherSuite: suite, Pattern: noise.HandshakeNN}
|
||||||
|
cfg.Initiator = true
|
||||||
|
hsI, err := noise.NewHandshakeState(cfg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
cfg.Initiator = false
|
||||||
|
hsR, err := noise.NewHandshakeState(cfg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
msg, _, _, err := hsI.WriteMessage(nil, nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
_, _, _, err = hsR.ReadMessage(nil, msg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
msg, dR, _, err := hsR.WriteMessage(nil, nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
_, eI, _, err := hsI.ReadMessage(nil, msg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, eI)
|
||||||
|
require.NotNil(t, dR)
|
||||||
|
|
||||||
|
// noise returns (cs1, cs2) where cs1 is the initiator->responder cipher.
|
||||||
|
return eI, dR
|
||||||
|
}
|
||||||
|
|
||||||
|
func roundtrip(t *testing.T, enc, dec CipherState) {
|
||||||
|
t.Helper()
|
||||||
|
plaintext := []byte("nebula cipher state roundtrip")
|
||||||
|
ad := []byte("aad")
|
||||||
|
nb := make([]byte, 12)
|
||||||
|
|
||||||
|
ct, err := enc.EncryptDanger(nil, ad, plaintext, 1, nb)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NotEqual(t, plaintext, ct)
|
||||||
|
|
||||||
|
pt, err := dec.DecryptDanger(nil, ad, ct, 1, nb)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, plaintext, pt)
|
||||||
|
|
||||||
|
// Wrong nonce must fail authentication.
|
||||||
|
_, err = dec.DecryptDanger(nil, ad, ct, 2, nb)
|
||||||
|
require.Error(t, err)
|
||||||
|
|
||||||
|
assert.Equal(t, enc.Overhead(), dec.Overhead())
|
||||||
|
assert.Equal(t, 16, enc.Overhead())
|
||||||
|
}
|
||||||
|
|
||||||
|
func BenchmarkCipherStateEncryptAESGCM(b *testing.B) {
|
||||||
|
enc, _ := buildCipherStatesB(b, CipherAESGCM)
|
||||||
|
benchEncryptCipherState(b, NewCipherState(enc, CipherAESGCM))
|
||||||
|
}
|
||||||
|
|
||||||
|
func BenchmarkCipherStateEncryptChaChaPoly(b *testing.B) {
|
||||||
|
enc, _ := buildCipherStatesB(b, noise.CipherChaChaPoly)
|
||||||
|
benchEncryptCipherState(b, NewCipherState(enc, noise.CipherChaChaPoly))
|
||||||
|
}
|
||||||
|
|
||||||
|
func benchEncryptCipherState(b *testing.B, cs CipherState) {
|
||||||
|
plaintext := make([]byte, 1280)
|
||||||
|
ad := make([]byte, 16)
|
||||||
|
nb := make([]byte, 12)
|
||||||
|
out := make([]byte, 0, len(plaintext)+cs.Overhead())
|
||||||
|
b.ResetTimer()
|
||||||
|
b.ReportAllocs()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
var err error
|
||||||
|
out, err = cs.EncryptDanger(out[:0], ad, plaintext, uint64(i+1), nb)
|
||||||
|
if err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildCipherStatesB(b *testing.B, c noise.CipherFunc) (*noise.CipherState, *noise.CipherState) {
|
||||||
|
b.Helper()
|
||||||
|
suite := noise.NewCipherSuite(noise.DH25519, c, noise.HashSHA256)
|
||||||
|
cfg := noise.Config{CipherSuite: suite, Pattern: noise.HandshakeNN}
|
||||||
|
cfg.Initiator = true
|
||||||
|
hsI, err := noise.NewHandshakeState(cfg)
|
||||||
|
if err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
cfg.Initiator = false
|
||||||
|
hsR, err := noise.NewHandshakeState(cfg)
|
||||||
|
if err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
msg, _, _, err := hsI.WriteMessage(nil, nil)
|
||||||
|
if err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, _, _, err := hsR.ReadMessage(nil, msg); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
msg, dR, _, err := hsR.WriteMessage(nil, nil)
|
||||||
|
if err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
_, eI, _, err := hsI.ReadMessage(nil, msg)
|
||||||
|
if err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
return eI, dR
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCipherStateNilSafety(t *testing.T) {
|
||||||
|
var aes *CipherStateAESGCM
|
||||||
|
_, err := aes.EncryptDanger(nil, nil, nil, 0, make([]byte, 12))
|
||||||
|
require.Error(t, err)
|
||||||
|
out, err := aes.DecryptDanger(nil, nil, nil, 0, make([]byte, 12))
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Empty(t, out)
|
||||||
|
assert.Equal(t, 0, aes.Overhead())
|
||||||
|
|
||||||
|
var cc *CipherStateChaChaPoly
|
||||||
|
_, err = cc.EncryptDanger(nil, nil, nil, 0, make([]byte, 12))
|
||||||
|
require.Error(t, err)
|
||||||
|
out, err = cc.DecryptDanger(nil, nil, nil, 0, make([]byte, 12))
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Empty(t, out)
|
||||||
|
assert.Equal(t, 0, cc.Overhead())
|
||||||
|
}
|
||||||
+276
-234
@@ -13,6 +13,7 @@ import (
|
|||||||
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -20,23 +21,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, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||||
err := h.Parse(packet)
|
err := h.Parse(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
||||||
|
// 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 +68,198 @@ 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, out, packet, h, fwPacket, lhf, nb, q, localCache, meta)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
out, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||||
|
if err != nil {
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Failed to decrypt packet",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"header", h,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Roam before we respond
|
||||||
|
f.handleHostRoaming(hostinfo, via)
|
||||||
|
f.connectionManager.In(hostinfo)
|
||||||
|
|
||||||
switch h.Type {
|
switch h.Type {
|
||||||
case header.Message:
|
case header.Message:
|
||||||
if !f.handleEncrypted(ci, via, h) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
switch h.Subtype {
|
switch h.Subtype {
|
||||||
case header.MessageNone:
|
case header.MessageNone:
|
||||||
if !f.decryptToTun(hostinfo, h.MessageCounter, out, packet, fwPacket, nb, q, localCache) {
|
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache, meta)
|
||||||
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:
|
||||||
|
//recycle the input packet ciphertext as our output buffer
|
||||||
|
f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, packet)
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||||
|
// The entire body is sent as AD, not encrypted.
|
||||||
|
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
||||||
|
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
|
||||||
|
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
|
||||||
|
// which will gracefully fail in the DecryptDanger call.
|
||||||
|
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||||
|
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
||||||
|
var err error
|
||||||
|
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// Advance the replay window now that the frame is authenticated
|
||||||
|
if !hostinfo.ConnectionState.window.Update(f.l, h.MessageCounter) {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("Unexpected packet received", "from", via)
|
hostinfo.logger(f.l).Debug("dropping out of window relay packet", "header", h)
|
||||||
}
|
}
|
||||||
return
|
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)
|
f.handleHostRoaming(hostinfo, via)
|
||||||
|
// Track usage of both the HostInfo and the Relay for the received & authenticated packet
|
||||||
f.connectionManager.In(hostinfo)
|
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",
|
||||||
|
"relayRemoteIndex", 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, meta)
|
||||||
|
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,
|
||||||
|
"relayFrom", hostinfo.vpnAddrs[0],
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
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 //todo it would potentially be nice to batch these
|
||||||
|
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
||||||
|
case TerminalType:
|
||||||
|
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Unexpected targetRelay Type", "from", via, "relayType", targetRelay.Type)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
hostinfo.logger(f.l).Info("Unexpected target relay state",
|
||||||
|
"relayTo", relay.PeerAddr,
|
||||||
|
"relayFrom", hostinfo.vpnAddrs[0],
|
||||||
|
"targetRelayState", targetRelay.State,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Unexpected relay type", "from", via, "relayType", relay.Type)
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// closeTunnel closes a tunnel locally, it does not send a closeTunnel packet to the remote
|
// closeTunnel closes a tunnel locally, it does not send a closeTunnel packet to the remote
|
||||||
@@ -270,7 +277,8 @@ func (f *Interface) sendCloseTunnel(h *HostInfo) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
|
func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
|
||||||
if !via.IsRelayed && hostinfo.remote != via.UdpAddr {
|
curRemote := hostinfo.GetRemote()
|
||||||
|
if !via.IsRelayed && curRemote != via.UdpAddr {
|
||||||
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
|
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("lighthouse.remote_allow_list denied roaming", "newAddr", via.UdpAddr)
|
hostinfo.logger(f.l).Debug("lighthouse.remote_allow_list denied roaming", "newAddr", via.UdpAddr)
|
||||||
@@ -282,7 +290,7 @@ func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
|
|||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("Suppressing roam back to previous remote",
|
hostinfo.logger(f.l).Debug("Suppressing roam back to previous remote",
|
||||||
"suppressSeconds", RoamingSuppressSeconds,
|
"suppressSeconds", RoamingSuppressSeconds,
|
||||||
"udpAddr", hostinfo.remote,
|
"udpAddr", curRemote,
|
||||||
"newAddr", via.UdpAddr,
|
"newAddr", via.UdpAddr,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -290,33 +298,16 @@ func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
hostinfo.logger(f.l).Info("Host roamed to new udp ip/port.",
|
hostinfo.logger(f.l).Info("Host roamed to new udp ip/port.",
|
||||||
"udpAddr", hostinfo.remote,
|
"udpAddr", curRemote,
|
||||||
"newAddr", via.UdpAddr,
|
"newAddr", via.UdpAddr,
|
||||||
)
|
)
|
||||||
hostinfo.lastRoam = time.Now()
|
hostinfo.lastRoam = time.Now()
|
||||||
hostinfo.lastRoamRemote = hostinfo.remote
|
hostinfo.lastRoamRemote = curRemote
|
||||||
hostinfo.SetRemote(via.UdpAddr)
|
hostinfo.SetRemote(via.UdpAddr)
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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")
|
||||||
@@ -431,16 +422,14 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
if dataLen <= offset+1 {
|
if dataLen <= offset+1 {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
next = (int(data[offset+1]) + 2) << 2
|
||||||
next = int(data[offset+1]+2) << 2
|
|
||||||
|
|
||||||
default:
|
default:
|
||||||
// Normal ipv6 header length processing
|
// Normal ipv6 header length processing
|
||||||
if dataLen <= offset+1 {
|
if dataLen <= offset+1 {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
next = (int(data[offset+1]) + 1) << 3
|
||||||
next = int(data[offset+1]+1) << 3
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if next <= 0 {
|
if next <= 0 {
|
||||||
@@ -523,60 +512,112 @@ 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 {
|
// 2-bit IP-level ECN codepoints (lower bits of IPv4 ToS / IPv6 TC).
|
||||||
var err error
|
const (
|
||||||
|
ecnNotECT = 0x00
|
||||||
|
ecnECT1 = 0x01
|
||||||
|
ecnECT0 = 0x02
|
||||||
|
ecnCE = 0x03
|
||||||
|
)
|
||||||
|
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
// applyOuterECN folds an outer CE mark from the underlay into the inner
|
||||||
if err != nil {
|
// IP header per RFC 6040 normal mode. It mutates pkt[1] in place. Other
|
||||||
hostinfo.logger(f.l).Error("Failed to decrypt packet", "error", err)
|
// codepoints are advisory only and leave the inner unchanged.
|
||||||
return false
|
//
|
||||||
|
// Merge cases (outer × inner → action):
|
||||||
|
//
|
||||||
|
// outer != CE : no-op (inner is authoritative)
|
||||||
|
// outer == CE, inner Not-ECT : log; cannot propagate to a non-ECN host
|
||||||
|
// outer == CE, inner ECT/CE : rewrite inner ECN to CE
|
||||||
|
func applyOuterECN(pkt []byte, outerECN byte, hostinfo *HostInfo, l *slog.Logger) {
|
||||||
|
if outerECN&ecnCE != ecnCE || len(pkt) < 2 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
switch pkt[0] >> 4 {
|
||||||
|
case 4:
|
||||||
|
switch pkt[1] & 0x03 {
|
||||||
|
case ecnNotECT:
|
||||||
|
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
||||||
|
}
|
||||||
|
case ecnCE:
|
||||||
|
// Already CE.
|
||||||
|
default:
|
||||||
|
// Rewriting the ToS byte invalidates the IPv4 header checksum, so
|
||||||
|
// patch it incrementally per RFC 1624 (HC' = ~(~HC + ~m + m')). The
|
||||||
|
// ToS is the low byte of the 16-bit word at pkt[0:2]; the header
|
||||||
|
// checksum lives at pkt[10:12]. A header too short to carry a
|
||||||
|
// checksum can't be fixed up here, so leave it for newPacket to
|
||||||
|
// reject rather than emit a mangled packet.
|
||||||
|
if len(pkt) < ipv4.HeaderLen {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
m := binary.BigEndian.Uint16(pkt[0:2])
|
||||||
|
pkt[1] = (pkt[1] &^ 0x03) | ecnCE
|
||||||
|
mNew := binary.BigEndian.Uint16(pkt[0:2])
|
||||||
|
sum := uint32(^binary.BigEndian.Uint16(pkt[10:12])) + uint32(^m) + uint32(mNew)
|
||||||
|
for sum > 0xffff {
|
||||||
|
sum = (sum >> 16) + (sum & 0xffff)
|
||||||
|
}
|
||||||
|
binary.BigEndian.PutUint16(pkt[10:12], ^uint16(sum))
|
||||||
|
}
|
||||||
|
case 6:
|
||||||
|
switch (pkt[1] >> 4) & 0x03 {
|
||||||
|
case ecnNotECT:
|
||||||
|
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
||||||
|
}
|
||||||
|
case ecnCE:
|
||||||
|
// Already CE.
|
||||||
|
default:
|
||||||
|
pkt[1] = (pkt[1] &^ 0x30) | (ecnCE << 4)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||||
|
// RFC 6040 normal-mode combine: fold any outer CE mark stamped by the
|
||||||
|
// underlay into the inner header before firewall + TUN write. Other
|
||||||
|
// outer codepoints are advisory only — we keep the inner unchanged.
|
||||||
|
if f.ecnEnabled.Load() {
|
||||||
|
applyOuterECN(out, meta.OuterECN, hostinfo, f.l)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = newPacket(out, true, fwPacket)
|
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)
|
||||||
if dropReason != nil {
|
if dropReason != nil {
|
||||||
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
|
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
|
||||||
// This gives us a buffer to build the reject packet in
|
// This gives us a buffer to build the reject packet in. With UDP GRO this is a single segment of a shared
|
||||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet, q)
|
// recvmmsg row whose capacity runs to the end of the whole row, so cap it to its own length (cap==len) to
|
||||||
|
// keep the reject builder from writing past this segment into the next, not-yet-processed coalesced segment.
|
||||||
|
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet[:len(packet):len(packet)], q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
||||||
"fwPacket", fwPacket,
|
"fwPacket", fwPacket,
|
||||||
"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) {
|
||||||
@@ -620,10 +661,11 @@ func (f *Interface) handleRecvError(addr netip.AddrPort, h *header.H) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if hostinfo.remote.IsValid() && hostinfo.remote != addr {
|
hr := hostinfo.GetRemote()
|
||||||
|
if hr.IsValid() && hr != addr {
|
||||||
f.l.Info("Someone spoofing recv_errors?",
|
f.l.Info("Someone spoofing recv_errors?",
|
||||||
"addr", addr,
|
"addr", addr,
|
||||||
"hostinfoRemote", hostinfo.remote,
|
"hostinfoRemote", hr,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -640,3 +640,38 @@ func serializeAH(ah *layers.IPSecAH) []byte {
|
|||||||
|
|
||||||
return buf.Bytes()
|
return buf.Bytes()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Test_newPacket_v6ExtHeaderOverflow is a regression test for the IPv6 extension-header
|
||||||
|
// length uint8 overflow in parseV6. A Destination-Options header with HdrExtLen=255 spans
|
||||||
|
// (255+1)*8 = 2048 bytes, so the real transport header sits at offset 2088. Before the fix
|
||||||
|
// the advance was computed in uint8 and wrapped to 0 (then clamped to 8), so the firewall
|
||||||
|
// read the transport header ~2KB too early from attacker-controlled option bytes while the
|
||||||
|
// host OS parses the real header, a firewall port/proto bypass. The fix makes parseV6 land
|
||||||
|
// on the same offset the host does.
|
||||||
|
func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
||||||
|
p := &firewall.Packet{}
|
||||||
|
|
||||||
|
const (
|
||||||
|
hdrLen = 40 // IPv6 header
|
||||||
|
extLen = 2048 // (255+1)*8, the true Destination-Options header size
|
||||||
|
realTCPAt = hdrLen + extLen // 2088, where the host reads the transport header
|
||||||
|
forgedTCPAt = hdrLen + 8 // 48, where the pre-fix wrapped+clamped walk landed
|
||||||
|
)
|
||||||
|
|
||||||
|
pkt := make([]byte, realTCPAt+4)
|
||||||
|
pkt[0] = 0x60 // version 6
|
||||||
|
pkt[6] = byte(layers.IPProtocolIPv6Destination) // NextHeader -> Destination Options
|
||||||
|
pkt[40] = byte(firewall.ProtoTCP) // Dest-Options NextHeader -> TCP
|
||||||
|
pkt[41] = 255 // HdrExtLen = 255
|
||||||
|
|
||||||
|
// Forged transport header at the pre-fix (wrong) offset: dst port 443.
|
||||||
|
binary.BigEndian.PutUint16(pkt[forgedTCPAt+2:forgedTCPAt+4], 443)
|
||||||
|
// Real transport header at the offset the host actually uses: dst port 22.
|
||||||
|
binary.BigEndian.PutUint16(pkt[realTCPAt+2:realTCPAt+4], 22)
|
||||||
|
|
||||||
|
require.NoError(t, newPacket(pkt, true, p))
|
||||||
|
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||||
|
// LocalPort is the destination port for incoming traffic. It must be the real port (22)
|
||||||
|
// the host delivers to, not the forged 443 at the overflowed offset.
|
||||||
|
assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset")
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,28 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import "net/netip"
|
||||||
|
|
||||||
|
type RxBatcher interface {
|
||||||
|
// Reserve creates a pkt to borrow
|
||||||
|
Reserve(sz int) []byte
|
||||||
|
// Commit borrows pkt. The caller must keep pkt valid until the next Flush
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
type TxBatcher interface {
|
||||||
|
// Reserve creates a pkt to borrow
|
||||||
|
Reserve(sz int) []byte
|
||||||
|
// Commit borrows pkt and records its destination plus the 2-bit
|
||||||
|
// IP-level ECN codepoint to set on the outer (carrier) header. The
|
||||||
|
// caller must keep pkt valid until the next Flush. Pass 0 (Not-ECT)
|
||||||
|
// to leave the outer ECN field unset.
|
||||||
|
Commit(pkt []byte, dst netip.AddrPort, outerECN byte)
|
||||||
|
// Flush emits every queued packet via the underlying batch writer in arrival order.
|
||||||
|
// Returns an errors.Join of one or more errors.
|
||||||
|
// After Flush returns, borrowed payload slices may be recycled.
|
||||||
|
Flush() error
|
||||||
|
}
|
||||||
@@ -0,0 +1,181 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/binary"
|
||||||
|
)
|
||||||
|
|
||||||
|
// flowKey identifies a transport flow by {src, dst, sport, dport, family}.
|
||||||
|
// Comparable, so map lookups and linear scans over the slot list stay tight.
|
||||||
|
// Shared by the TCP and UDP coalescers; each coalescer keeps its own
|
||||||
|
// openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto
|
||||||
|
// never alias.
|
||||||
|
type flowKey struct {
|
||||||
|
src, dst [16]byte
|
||||||
|
sport, dport uint16
|
||||||
|
isV6 bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// initialSlots is the starting capacity of the slot pool. One flow per
|
||||||
|
// packet is the worst case so this matches a typical carrier-side
|
||||||
|
// recvmmsg batch on the encrypted UDP socket.
|
||||||
|
const initialSlots = 64
|
||||||
|
|
||||||
|
// parsedIP is the IP-level result of parseIPPrologue. The caller layers
|
||||||
|
// L4-specific parsing (TCP / UDP) on top.
|
||||||
|
type parsedIP struct {
|
||||||
|
fk flowKey
|
||||||
|
ipHdrLen int
|
||||||
|
// pkt is the original buffer trimmed to the IP-declared total length.
|
||||||
|
// Anything below the IP layer (transport parsers) should slice into
|
||||||
|
// pkt rather than the unbounded original.
|
||||||
|
pkt []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseIPPrologue extracts the IP-level fields the coalescers care about:
|
||||||
|
// IHL/payload length, version, src/dst addresses, and the L4 protocol byte.
|
||||||
|
// Returns ok=false for malformed input, IPv4 with options or fragmentation,
|
||||||
|
// or IPv6 with extension headers (all rejected by both coalescers in
|
||||||
|
// identical ways before this refactor).
|
||||||
|
//
|
||||||
|
// On success, p.pkt is len-trimmed to the IP-declared length so callers
|
||||||
|
// don't have to repeat the trim. wantProto is the IANA protocol number to
|
||||||
|
// require (6 for TCP, 17 for UDP); ok=false for any other value.
|
||||||
|
func parseIPPrologue(pkt []byte, wantProto byte) (parsedIP, bool) {
|
||||||
|
var p parsedIP
|
||||||
|
if len(pkt) < 20 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
v := pkt[0] >> 4
|
||||||
|
switch v {
|
||||||
|
case 4:
|
||||||
|
ihl := int(pkt[0]&0x0f) * 4
|
||||||
|
if ihl != 20 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
if pkt[9] != wantProto {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
// Reject actual fragmentation (MF or non-zero frag offset).
|
||||||
|
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
|
||||||
|
if totalLen > len(pkt) || totalLen < ihl {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.ipHdrLen = 20
|
||||||
|
p.fk.isV6 = false
|
||||||
|
copy(p.fk.src[:4], pkt[12:16])
|
||||||
|
copy(p.fk.dst[:4], pkt[16:20])
|
||||||
|
p.pkt = pkt[:totalLen]
|
||||||
|
case 6:
|
||||||
|
if len(pkt) < 40 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
if pkt[6] != wantProto {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
|
||||||
|
if 40+payloadLen > len(pkt) {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.ipHdrLen = 40
|
||||||
|
p.fk.isV6 = true
|
||||||
|
copy(p.fk.src[:], pkt[8:24])
|
||||||
|
copy(p.fk.dst[:], pkt[24:40])
|
||||||
|
p.pkt = pkt[:40+payloadLen]
|
||||||
|
default:
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
return p, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// ipHeadersMatch compares the IP portion of two packet header prefixes for
|
||||||
|
// byte-for-byte equality on every field that must be identical across
|
||||||
|
// coalesced segments. Size/IPID/IPCsum are masked out. The full DSCP/ECN
|
||||||
|
// byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel
|
||||||
|
// GRO: segments with differing ECN codepoints must not coalesce, otherwise
|
||||||
|
// ORing e.g. ECT(0) with ECT(1) would fabricate a false CE (congestion)
|
||||||
|
// mark or mark a Not-ECT flow as ECN-capable.
|
||||||
|
//
|
||||||
|
// The transport (L4) portion of the header is checked separately by the
|
||||||
|
// per-protocol matcher.
|
||||||
|
func ipHeadersMatch(a, b []byte, isV6 bool) bool {
|
||||||
|
if isV6 {
|
||||||
|
// IPv6: byte 0 = version/TC[7:4], byte 1 = TC[3:0]/flow[19:16],
|
||||||
|
// bytes [2:4] = flow[15:0], [6:8] = next_hdr/hop, [8:40] = src+dst.
|
||||||
|
// Compare byte 1 fully so ECN (TC[1:0]) must match. Skip [4:6] payload_len.
|
||||||
|
if a[0] != b[0] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if a[1] != b[1] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[2:4], b[2:4]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[6:40], b[6:40]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
// IPv4: byte 0 = version/IHL, byte 1 = DSCP(6)|ECN(2),
|
||||||
|
// [6:10] flags/fragoff/TTL/proto, [12:20] src+dst.
|
||||||
|
// Compare byte 1 fully so ECN must match.
|
||||||
|
// Skip [2:4] total len, [4:6] id, [10:12] csum.
|
||||||
|
if a[0] != b[0] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if a[1] != b[1] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[6:10], b[6:10]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[12:20], b[12:20]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
|
||||||
|
// slices via Reserve and releases them in bulk via Reset.
|
||||||
|
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]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reserver hands out an sz-byte slice valid until its Resetter runs.
|
||||||
|
type Reserver func(sz int) []byte
|
||||||
|
|
||||||
|
// Resetter clears all reservations held by a Reserver. Only the arena's
|
||||||
|
// owner holds one; lanes inside a MultiCoalescer get nil.
|
||||||
|
type Resetter func()
|
||||||
@@ -0,0 +1,132 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// MultiCoalescer fans plaintext packets out to lane-specific batchers based
|
||||||
|
// on the IP/L4 protocol of the packet, sharing a single Reserve arena
|
||||||
|
// across lanes so the caller's allocation pattern is unchanged.
|
||||||
|
//
|
||||||
|
// Lanes are processed independently: the TCP coalescer only sees TCP, the
|
||||||
|
// UDP coalescer only sees UDP, and the passthrough lane handles everything
|
||||||
|
// else. Per-flow arrival order is preserved because a single 5-tuple only
|
||||||
|
// ever lands in one lane and each lane preserves its own slot order.
|
||||||
|
//
|
||||||
|
// Cross-lane order is NOT preserved across the TCP/UDP/passthrough split.
|
||||||
|
// This is acceptable because the carrier-side recvmmsg path already
|
||||||
|
// stable-sorts by (peer, message counter) before delivering plaintext
|
||||||
|
// here, so replay-window invariants are unaffected, and apps observe
|
||||||
|
// correct per-flow ordering — which is all the IP layer guarantees anyway.
|
||||||
|
// Do not "fix" this by interleaving lane outputs at flush time; that
|
||||||
|
// negates the entire point of coalescing (each lane needs to see runs of
|
||||||
|
// adjacent same-flow packets to coalesce them).
|
||||||
|
type MultiCoalescer struct {
|
||||||
|
tcp *TCPCoalescer
|
||||||
|
udp *UDPCoalescer
|
||||||
|
pt *Passthrough
|
||||||
|
// arena is owned by the Multi: lanes get only its Reserve (nil Resetter)
|
||||||
|
// and Flush resets it exactly once after every lane has drained.
|
||||||
|
arena *Arena
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultMultiArenaCap is the recommended arena capacity for a Multi-lane
|
||||||
|
// batcher: 64 slots × 65535 bytes ≈ 4 MiB, enough to hold one recvmmsg
|
||||||
|
// burst worth of MTU-sized packets without the arena growing.
|
||||||
|
const DefaultMultiArenaCap = initialSlots * 65535
|
||||||
|
|
||||||
|
// NewMultiCoalescer builds a multi-lane batcher. tcpEnabled lets the caller
|
||||||
|
// opt out of TCP coalescing (e.g. when the queue can't do TSO); udpEnabled
|
||||||
|
// likewise gates UDP coalescing (only enable when USO was negotiated).
|
||||||
|
// Either lane disabled redirects its traffic into the passthrough lane.
|
||||||
|
// arena is the single backing slab shared across every lane; the caller
|
||||||
|
// pre-sizes it via NewArena so the hot path never allocates.
|
||||||
|
func NewMultiCoalescer(w io.Writer, l *slog.Logger, arena *Arena, tcpEnabled, udpEnabled bool) *MultiCoalescer {
|
||||||
|
m := &MultiCoalescer{
|
||||||
|
pt: NewPassthrough(w, arena.Reserve, nil),
|
||||||
|
arena: arena,
|
||||||
|
}
|
||||||
|
if tcpEnabled {
|
||||||
|
m.tcp = NewTCPCoalescer(w, l, arena.Reserve, nil)
|
||||||
|
}
|
||||||
|
if udpEnabled {
|
||||||
|
m.udp = NewUDPCoalescer(w, arena.Reserve, nil)
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *MultiCoalescer) Reserve(sz int) []byte {
|
||||||
|
return m.arena.Reserve(sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Commit dispatches pkt to the appropriate lane based on IP version + L4
|
||||||
|
// proto. Borrowed slice contract is identical to the single-lane batchers,
|
||||||
|
// pkt must remain valid until the next Flush.
|
||||||
|
//
|
||||||
|
// On the success path the IP/TCP-or-UDP parse happens here once and the
|
||||||
|
// parsed struct is handed to the lane via commitParsed so the lane doesn't
|
||||||
|
// re-walk the header.
|
||||||
|
func (m *MultiCoalescer) Commit(pkt []byte) error {
|
||||||
|
if len(pkt) < 20 {
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
v := pkt[0] >> 4
|
||||||
|
var proto byte
|
||||||
|
switch v {
|
||||||
|
case 4:
|
||||||
|
proto = pkt[9]
|
||||||
|
case 6:
|
||||||
|
if len(pkt) < 40 {
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
proto = pkt[6]
|
||||||
|
default:
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
switch proto {
|
||||||
|
case ipProtoTCP:
|
||||||
|
if m.tcp != nil {
|
||||||
|
info, ok := parseTCPBase(pkt)
|
||||||
|
if !ok {
|
||||||
|
// Malformed/unsupported TCP shape (IP options, fragments, ...).
|
||||||
|
// Handle this via passthrough support in the TCP coalescer, to attempt to preserve flow order.
|
||||||
|
m.tcp.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return m.tcp.commitParsed(pkt, info)
|
||||||
|
}
|
||||||
|
case ipProtoUDP:
|
||||||
|
if m.udp != nil {
|
||||||
|
info, ok := parseUDP(pkt)
|
||||||
|
if !ok {
|
||||||
|
m.udp.addPassthrough(pkt) //we could also m.pt.Commit() here I guess?
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return m.udp.commitParsed(pkt, info)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush drains every lane in a fixed order, then resets the shared arena once.
|
||||||
|
// A lane error doesn't stop the remaining lanes; the joined errors are returned.
|
||||||
|
func (m *MultiCoalescer) Flush() error {
|
||||||
|
var errs []error
|
||||||
|
if m.tcp != nil {
|
||||||
|
if err := m.tcp.Flush(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if m.udp != nil {
|
||||||
|
if err := m.udp.Flush(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := m.pt.Flush(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
m.arena.Reset()
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
@@ -0,0 +1,96 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right
|
||||||
|
// lane: TCP and UDP get coalesced when their lanes are enabled, anything
|
||||||
|
// else (ICMP here) falls through to plain Write.
|
||||||
|
func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, true)
|
||||||
|
|
||||||
|
tcpPay := make([]byte, 1200)
|
||||||
|
udpPay := make([]byte, 1200)
|
||||||
|
icmp := make([]byte, 28)
|
||||||
|
icmp[0] = 0x45
|
||||||
|
icmp[2] = 0
|
||||||
|
icmp[3] = 28
|
||||||
|
icmp[9] = 1
|
||||||
|
|
||||||
|
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(icmp); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
// 1 TCP super (2 segments) + 1 UDP super (2 segments) = 2 gso writes.
|
||||||
|
if len(w.gsoWrites) != 2 {
|
||||||
|
t.Fatalf("want 2 gso writes (one TCP + one UDP), got %d", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
if len(w.writes) != 1 {
|
||||||
|
t.Fatalf("want 1 plain write (ICMP), got %d", len(w.writes))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestMultiCoalescerDisabledUDPFallsThrough verifies that when the UDP lane
|
||||||
|
// is disabled (e.g. kernel doesn't support USO), UDP packets still reach
|
||||||
|
// the kernel via the passthrough lane rather than being lost.
|
||||||
|
func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, false) // TSO on, USO off
|
||||||
|
|
||||||
|
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 0 {
|
||||||
|
t.Errorf("UDP must NOT be coalesced when USO disabled, got %d gso writes", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
if len(w.writes) != 2 {
|
||||||
|
t.Errorf("UDP must pass through as 2 plain writes, got %d", len(w.writes))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestMultiCoalescerDisabledTCPFallsThrough mirrors the TSO=off case.
|
||||||
|
func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), false, true) // TSO off, USO on
|
||||||
|
|
||||||
|
pay := make([]byte, 1200)
|
||||||
|
if err := m.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 0 {
|
||||||
|
t.Errorf("TCP must NOT be coalesced when TSO disabled, got %d gso writes", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
if len(w.writes) != 2 {
|
||||||
|
t.Errorf("TCP must pass through as 2 plain writes, got %d", len(w.writes))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Passthrough is a RxBatcher that doesn't batch anything, it just accumulates and then sends packets.
|
||||||
|
type Passthrough struct {
|
||||||
|
out io.Writer
|
||||||
|
slots [][]byte
|
||||||
|
reserver Reserver
|
||||||
|
resetter Resetter
|
||||||
|
cursor int
|
||||||
|
}
|
||||||
|
|
||||||
|
const passthroughBaseNumSlots = 128
|
||||||
|
|
||||||
|
// DefaultPassthroughArenaCap is the recommended arena capacity for a
|
||||||
|
// standalone Passthrough batcher: 128 slots × udp.MTU ≈ 1.1 MiB.
|
||||||
|
const DefaultPassthroughArenaCap = passthroughBaseNumSlots * udp.MTU
|
||||||
|
|
||||||
|
func NewPassthrough(w io.Writer, reserver Reserver, resetter Resetter) *Passthrough {
|
||||||
|
return &Passthrough{
|
||||||
|
out: w,
|
||||||
|
slots: make([][]byte, 0, passthroughBaseNumSlots),
|
||||||
|
reserver: reserver,
|
||||||
|
resetter: resetter,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Reserve(sz int) []byte {
|
||||||
|
return p.reserver(sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Commit(pkt []byte) error {
|
||||||
|
p.slots = append(p.slots, pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush drains every queued packet and calls the configured Resetter
|
||||||
|
func (p *Passthrough) Flush() error {
|
||||||
|
firstErr := p.drain()
|
||||||
|
if p.resetter != nil {
|
||||||
|
p.resetter()
|
||||||
|
}
|
||||||
|
return firstErr
|
||||||
|
}
|
||||||
|
|
||||||
|
// drain writes out every queued packet and clears the slot list.
|
||||||
|
func (p *Passthrough) drain() 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]
|
||||||
|
return firstErr
|
||||||
|
}
|
||||||
@@ -0,0 +1,728 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/binary"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ipProtoTCP is the IANA protocol number for TCP. Defined here to help Windows out.
|
||||||
|
const ipProtoTCP = 6
|
||||||
|
|
||||||
|
// tcpCoalesceBufSize caps total bytes per superpacket. Mirrors the kernel's
|
||||||
|
// sk_gso_max_size of ~64KiB; anything beyond this would be rejected anyway.
|
||||||
|
const tcpCoalesceBufSize = 65535
|
||||||
|
|
||||||
|
// tcpCoalesceMaxSegs caps how many segments we'll coalesce into a single
|
||||||
|
// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency.
|
||||||
|
const tcpCoalesceMaxSegs = 64
|
||||||
|
|
||||||
|
// tcpCoalesceHdrCap is the scratch space we copy a seed's IP+TCP header
|
||||||
|
// into. IPv6 (40) + TCP with full options (60) = 100 bytes.
|
||||||
|
const tcpCoalesceHdrCap = 100
|
||||||
|
|
||||||
|
// coalesceSlot is one entry in the coalescer's ordered event queue. When
|
||||||
|
// passthrough is true the slot holds a single borrowed packet that must be
|
||||||
|
// emitted verbatim (non-TCP, non-admissible TCP, or oversize seed). When
|
||||||
|
// passthrough is false the slot is an in-progress coalesced superpacket:
|
||||||
|
// hdrBuf is a mutable copy of the seed's IP+TCP header (we patch total
|
||||||
|
// length and pseudo-header partial at flush), and payIovs are *borrowed*
|
||||||
|
// slices from the caller's plaintext buffers — no payload is ever copied.
|
||||||
|
// The caller (listenOut) must keep those buffers alive until Flush.
|
||||||
|
type coalesceSlot struct {
|
||||||
|
passthrough bool
|
||||||
|
rawPkt []byte // borrowed when passthrough
|
||||||
|
|
||||||
|
fk flowKey
|
||||||
|
hdrBuf [tcpCoalesceHdrCap]byte
|
||||||
|
hdrLen int
|
||||||
|
ipHdrLen int
|
||||||
|
isV6 bool
|
||||||
|
gsoSize int
|
||||||
|
numSeg int
|
||||||
|
totalPay int
|
||||||
|
nextSeq uint32
|
||||||
|
// psh closes the chain: set when the last-accepted segment had PSH or
|
||||||
|
// was sub-gsoSize. No further appends after that.
|
||||||
|
psh bool
|
||||||
|
payIovs [][]byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// TCPCoalescer accumulates adjacent in-flow TCP data segments across
|
||||||
|
// multiple concurrent flows and emits each flow's run as a single TSO
|
||||||
|
// superpacket via tio.GSOWriter. All output — coalesced or not — is
|
||||||
|
// deferred until Flush so arrival order is preserved on the wire. Owns
|
||||||
|
// no locks; one coalescer per TUN write queue.
|
||||||
|
type TCPCoalescer struct {
|
||||||
|
plainW io.Writer
|
||||||
|
gsoW tio.GSOWriter // nil when the queue doesn't support TSO
|
||||||
|
|
||||||
|
// slots is the ordered event queue. Flush walks it once and emits each
|
||||||
|
// entry as either a WriteGSO (coalesced) or a plainW.Write (passthrough).
|
||||||
|
slots []*coalesceSlot
|
||||||
|
// openSlots maps a flow key to its most recent non-sealed slot, so new
|
||||||
|
// segments can extend an in-progress superpacket in O(1). Slots are
|
||||||
|
// removed from this map when they close (PSH or short-last-segment),
|
||||||
|
// when a non-admissible packet for that flow arrives, or in Flush.
|
||||||
|
openSlots map[flowKey]*coalesceSlot
|
||||||
|
// lastSlot caches the most recently touched open slot. Steady-state
|
||||||
|
// bulk traffic is dominated by a single flow, so comparing the
|
||||||
|
// incoming key against the cached slot's own fk lets the hot path
|
||||||
|
// skip the map lookup (and the aeshash of a 38-byte key) entirely.
|
||||||
|
// Kept in lockstep with openSlots: nil whenever the slot it pointed
|
||||||
|
// at is removed/sealed.
|
||||||
|
lastSlot *coalesceSlot
|
||||||
|
pool []*coalesceSlot // free list for reuse
|
||||||
|
reserver Reserver
|
||||||
|
resetter Resetter
|
||||||
|
l *slog.Logger
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewTCPCoalescer(w io.Writer, l *slog.Logger, reserver Reserver, resetter Resetter) *TCPCoalescer {
|
||||||
|
c := &TCPCoalescer{
|
||||||
|
plainW: w,
|
||||||
|
slots: make([]*coalesceSlot, 0, initialSlots),
|
||||||
|
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
|
||||||
|
pool: make([]*coalesceSlot, 0, initialSlots),
|
||||||
|
reserver: reserver,
|
||||||
|
resetter: resetter,
|
||||||
|
l: l,
|
||||||
|
}
|
||||||
|
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP); ok {
|
||||||
|
c.gsoW = gw
|
||||||
|
}
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// parsedTCP holds the fields extracted from a single parse so later steps
|
||||||
|
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||||
|
type parsedTCP struct {
|
||||||
|
fk flowKey
|
||||||
|
ipHdrLen int
|
||||||
|
tcpHdrLen int
|
||||||
|
hdrLen int
|
||||||
|
payLen int
|
||||||
|
seq uint32
|
||||||
|
flags byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseTCPBase extracts the flow key and IP/TCP offsets for any TCP packet,
|
||||||
|
// regardless of whether it's admissible for coalescing. Returns ok=false
|
||||||
|
// for non-TCP or malformed input.
|
||||||
|
// Accepts IPv4 (no options or fragmentation) and IPv6 (no extension headers).
|
||||||
|
func parseTCPBase(pkt []byte) (parsedTCP, bool) {
|
||||||
|
var p parsedTCP
|
||||||
|
ip, ok := parseIPPrologue(pkt, ipProtoTCP)
|
||||||
|
if !ok {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
pkt = ip.pkt
|
||||||
|
p.fk = ip.fk
|
||||||
|
p.ipHdrLen = ip.ipHdrLen
|
||||||
|
|
||||||
|
if len(pkt) < p.ipHdrLen+20 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
tcpOff := int(pkt[p.ipHdrLen+12]>>4) * 4
|
||||||
|
if tcpOff < 20 || tcpOff > 60 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
if len(pkt) < p.ipHdrLen+tcpOff {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.tcpHdrLen = tcpOff
|
||||||
|
p.hdrLen = p.ipHdrLen + tcpOff
|
||||||
|
p.payLen = len(pkt) - p.hdrLen
|
||||||
|
p.seq = binary.BigEndian.Uint32(pkt[p.ipHdrLen+4 : p.ipHdrLen+8])
|
||||||
|
p.flags = pkt[p.ipHdrLen+13]
|
||||||
|
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
||||||
|
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
||||||
|
return p, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// TCP flag bits (byte 13 of the TCP header). Only the bits actually consulted
|
||||||
|
// by the coalescer are named; FIN/SYN/RST/URG/CWR are rejected via the
|
||||||
|
// negative mask in coalesceable, not by name.
|
||||||
|
const (
|
||||||
|
tcpFlagPsh = 0x08
|
||||||
|
tcpFlagAck = 0x10
|
||||||
|
tcpFlagEce = 0x40
|
||||||
|
)
|
||||||
|
|
||||||
|
// coalesceable reports whether a parsed TCP segment is eligible for
|
||||||
|
// coalescing. Accepts ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE with a
|
||||||
|
// non-empty payload. CWR is excluded because it marks a one-shot
|
||||||
|
// congestion-window-reduced transition the receiver must observe at a
|
||||||
|
// segment boundary.
|
||||||
|
func (p parsedTCP) coalesceable() bool {
|
||||||
|
if p.flags&tcpFlagAck == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if p.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return p.payLen > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) Reserve(sz int) []byte {
|
||||||
|
return c.reserver(sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||||
|
func (c *TCPCoalescer) Commit(pkt []byte) error {
|
||||||
|
if c.gsoW == nil {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
info, ok := parseTCPBase(pkt)
|
||||||
|
if !ok {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return c.commitParsed(pkt, info)
|
||||||
|
}
|
||||||
|
|
||||||
|
// commitParsed is the post-parse half of Commit. The caller must have
|
||||||
|
// already verified parseTCPBase succeeded (info is a valid TCP parse).
|
||||||
|
// Used by MultiCoalescer.Commit to avoid re-walking the IP/TCP header
|
||||||
|
// after the dispatcher has already done so.
|
||||||
|
func (c *TCPCoalescer) commitParsed(pkt []byte, info parsedTCP) error {
|
||||||
|
if c.gsoW == nil {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if !info.coalesceable() {
|
||||||
|
// TCP but not admissible (SYN/FIN/RST/URG/CWR or zero-payload).
|
||||||
|
// Seal this flow's open slot so later in-flow packets don't extend
|
||||||
|
// it and accidentally reorder past this passthrough.
|
||||||
|
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||||
|
c.lastSlot = nil
|
||||||
|
}
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Single-flow fast path: with only one open flow the cache hits every
|
||||||
|
// packet, and len(openSlots)==1 lets us skip the 38-byte fk compare
|
||||||
|
// when there are multiple flows in flight (where the hit rate would
|
||||||
|
// be ~0 and the compare is pure overhead).
|
||||||
|
var open *coalesceSlot
|
||||||
|
if last := c.lastSlot; last != nil && len(c.openSlots) == 1 && last.fk == info.fk {
|
||||||
|
open = last
|
||||||
|
} else {
|
||||||
|
open = c.openSlots[info.fk]
|
||||||
|
}
|
||||||
|
if open != nil {
|
||||||
|
if c.canAppend(open, pkt, info) {
|
||||||
|
c.appendPayload(open, pkt, info)
|
||||||
|
if open.psh {
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
c.lastSlot = nil
|
||||||
|
} else {
|
||||||
|
c.lastSlot = open
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// Can't extend — seal it and fall through to seed a fresh slot.
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
if c.lastSlot == open {
|
||||||
|
c.lastSlot = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
c.seed(pkt, info)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush emits every queued event in (per-flow) seq order.
|
||||||
|
func (c *TCPCoalescer) Flush() error {
|
||||||
|
first := c.drain()
|
||||||
|
if c.resetter != nil {
|
||||||
|
c.resetter()
|
||||||
|
}
|
||||||
|
return first
|
||||||
|
}
|
||||||
|
|
||||||
|
// drain emits every queued slot (reordering/merging coalesced runs first)
|
||||||
|
// and clears the slot state.
|
||||||
|
func (c *TCPCoalescer) drain() error {
|
||||||
|
c.reorderForFlush()
|
||||||
|
var first error
|
||||||
|
for _, s := range c.slots {
|
||||||
|
var err error
|
||||||
|
if s.passthrough {
|
||||||
|
_, err = c.plainW.Write(s.rawPkt)
|
||||||
|
} else {
|
||||||
|
err = c.flushSlot(s)
|
||||||
|
}
|
||||||
|
if err != nil && first == nil {
|
||||||
|
first = err
|
||||||
|
}
|
||||||
|
c.release(s)
|
||||||
|
}
|
||||||
|
clear(c.slots)
|
||||||
|
c.slots = c.slots[:0]
|
||||||
|
clear(c.openSlots)
|
||||||
|
c.lastSlot = nil
|
||||||
|
|
||||||
|
return first
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) addPassthrough(pkt []byte) {
|
||||||
|
s := c.take()
|
||||||
|
s.passthrough = true
|
||||||
|
s.rawPkt = pkt
|
||||||
|
c.slots = append(c.slots, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
|
||||||
|
if info.hdrLen > tcpCoalesceHdrCap || info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
||||||
|
// Pathological shape — can't fit our scratch, emit as-is.
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s := c.take()
|
||||||
|
s.passthrough = false
|
||||||
|
s.rawPkt = nil
|
||||||
|
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
||||||
|
s.hdrLen = info.hdrLen
|
||||||
|
s.ipHdrLen = info.ipHdrLen
|
||||||
|
s.isV6 = info.fk.isV6
|
||||||
|
s.fk = info.fk
|
||||||
|
s.gsoSize = info.payLen
|
||||||
|
s.numSeg = 1
|
||||||
|
s.totalPay = info.payLen
|
||||||
|
s.nextSeq = info.seq + uint32(info.payLen)
|
||||||
|
s.psh = info.flags&tcpFlagPsh != 0
|
||||||
|
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
|
c.slots = append(c.slots, s)
|
||||||
|
if !s.psh {
|
||||||
|
c.openSlots[info.fk] = s
|
||||||
|
c.lastSlot = s
|
||||||
|
} else if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||||
|
// PSH-on-seed seals the slot immediately. Any prior cached open
|
||||||
|
// slot for this flow has just been sealed-and-replaced by this
|
||||||
|
// passthrough-shaped seed, so drop the cache too.
|
||||||
|
c.lastSlot = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// canAppend reports whether info's packet extends the slot's seed: same
|
||||||
|
// header shape and stable contents, adjacent seq, not oversized, chain not closed.
|
||||||
|
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bool {
|
||||||
|
if s.psh {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if info.hdrLen != s.hdrLen {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if info.seq != s.nextSeq {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if s.numSeg >= tcpCoalesceMaxSegs {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if info.payLen > s.gsoSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
// ECE state must be stable across a burst — receivers expect the
|
||||||
|
// flag set on every segment of a CE-echoing window or none.
|
||||||
|
seedFlags := s.hdrBuf[s.ipHdrLen+13]
|
||||||
|
if (seedFlags^info.flags)&tcpFlagEce != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !headersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info parsedTCP) {
|
||||||
|
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
|
s.numSeg++
|
||||||
|
s.totalPay += info.payLen
|
||||||
|
s.nextSeq = info.seq + uint32(info.payLen)
|
||||||
|
if info.flags&tcpFlagPsh != 0 {
|
||||||
|
// Propagate PSH into the seed header so kernel TSO sets it on the
|
||||||
|
// last segment. Without this the sender's push signal is dropped.
|
||||||
|
s.hdrBuf[s.ipHdrLen+13] |= tcpFlagPsh
|
||||||
|
}
|
||||||
|
if info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0 {
|
||||||
|
s.psh = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) take() *coalesceSlot {
|
||||||
|
if n := len(c.pool); n > 0 {
|
||||||
|
s := c.pool[n-1]
|
||||||
|
c.pool[n-1] = nil
|
||||||
|
c.pool = c.pool[:n-1]
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
return &coalesceSlot{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
||||||
|
s.passthrough = false
|
||||||
|
s.rawPkt = nil
|
||||||
|
clear(s.payIovs)
|
||||||
|
s.payIovs = s.payIovs[:0]
|
||||||
|
s.numSeg = 0
|
||||||
|
s.totalPay = 0
|
||||||
|
s.psh = false
|
||||||
|
c.pool = append(c.pool, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
// flushSlot patches the header and calls WriteGSO. Does not remove the slot from c.slots.
|
||||||
|
func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
||||||
|
total := s.hdrLen + s.totalPay
|
||||||
|
l4Len := total - s.ipHdrLen
|
||||||
|
hdr := s.hdrBuf[:s.hdrLen]
|
||||||
|
|
||||||
|
if s.isV6 {
|
||||||
|
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
||||||
|
} else {
|
||||||
|
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
|
||||||
|
hdr[10] = 0
|
||||||
|
hdr[11] = 0
|
||||||
|
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
|
||||||
|
}
|
||||||
|
|
||||||
|
var psum uint32
|
||||||
|
if s.isV6 {
|
||||||
|
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoTCP, l4Len)
|
||||||
|
} else {
|
||||||
|
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoTCP, l4Len)
|
||||||
|
}
|
||||||
|
tcsum := s.ipHdrLen + 16
|
||||||
|
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
||||||
|
|
||||||
|
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
|
||||||
|
}
|
||||||
|
|
||||||
|
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
|
||||||
|
// equality on every field that must be identical across coalesced
|
||||||
|
// segments. Size/IPID/IPCsum/seq/flags/tcpCsum are masked out.
|
||||||
|
func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||||
|
if len(a) != len(b) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !ipHeadersMatch(a, b, isV6) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
// TCP: compare [0:4] ports, [8:13] ack+dataoff, [14:16] window,
|
||||||
|
// [18:tcpHdrLen] options (incl. urgent).
|
||||||
|
tcp := ipHdrLen
|
||||||
|
if !bytes.Equal(a[tcp:tcp+4], b[tcp:tcp+4]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[tcp+8:tcp+13], b[tcp+8:tcp+13]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[tcp+14:tcp+16], b[tcp+14:tcp+16]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[tcp+18:], b[tcp+18:]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// reorderForFlush neutralizes wire-side reorder that the rxOrder buffer
|
||||||
|
// couldn't catch (anything crossing a recvmmsg batch boundary). Without
|
||||||
|
// this pass a small wire reorder — counter 250 arriving in batch K when
|
||||||
|
// 200..249 are coming in batch K+1 — would seed an out-of-seq slot first
|
||||||
|
// and emit it ahead of the lower-seq slot, manifesting at the inner TCP
|
||||||
|
// receiver as a much larger reorder than the wire actually had.
|
||||||
|
//
|
||||||
|
// Two phases:
|
||||||
|
// 1. Sort each passthrough-bounded segment of c.slots by (flow, seq).
|
||||||
|
// Cross-flow ordering inside a segment isn't preserved (it never was
|
||||||
|
// and doesn't matter for any single flow's TCP correctness).
|
||||||
|
// 2. Sweep once and merge adjacent same-flow slots whose ranges are now
|
||||||
|
// contiguous AND whose tail is gsoSize-aligned. The tail constraint
|
||||||
|
// matters because the kernel TSO splitter chops at gsoSize from the
|
||||||
|
// start of the merged payload — a short segment in the middle would
|
||||||
|
// desynchronize every later segment.
|
||||||
|
//
|
||||||
|
// Passthrough slots act as barriers: the merge check skips them on either
|
||||||
|
// side, so a SYN/FIN/RST/CWR is never reordered relative to its flow's
|
||||||
|
// data.
|
||||||
|
func (c *TCPCoalescer) reorderForFlush() {
|
||||||
|
if len(c.slots) <= 1 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
runStart := 0
|
||||||
|
for i := 0; i <= len(c.slots); i++ {
|
||||||
|
if i < len(c.slots) && !c.slots[i].passthrough {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
c.sortRun(c.slots[runStart:i])
|
||||||
|
runStart = i + 1
|
||||||
|
}
|
||||||
|
out := c.slots[:0]
|
||||||
|
logged := false
|
||||||
|
for _, s := range c.slots {
|
||||||
|
if n := len(out); n > 0 {
|
||||||
|
prev := out[n-1]
|
||||||
|
if !prev.passthrough && !s.passthrough && prev.fk == s.fk {
|
||||||
|
// Same-flow neighbors after sort. If they aren't seq-
|
||||||
|
// contiguous it's a real gap — packets the wire reordered
|
||||||
|
// across batches, or actual loss before nebula. Log it so
|
||||||
|
// the operator can quantify how often it happens; the data
|
||||||
|
// itself still emits in seq order, kernel TCP handles the
|
||||||
|
// gap via its OOO queue.
|
||||||
|
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
if prev.nextSeq != slotSeedSeq(s) {
|
||||||
|
logged = true
|
||||||
|
gap := int64(slotSeedSeq(s)) - int64(prev.nextSeq)
|
||||||
|
c.l.Debug("tcp coalesce: cross-slot seq gap",
|
||||||
|
"src", flowKeyAddr(s.fk, false),
|
||||||
|
"dst", flowKeyAddr(s.fk, true),
|
||||||
|
"sport", s.fk.sport,
|
||||||
|
"dport", s.fk.dport,
|
||||||
|
"prev_seed_seq", slotSeedSeq(prev),
|
||||||
|
"prev_next_seq", prev.nextSeq,
|
||||||
|
"this_seed_seq", slotSeedSeq(s),
|
||||||
|
"gap_bytes", gap,
|
||||||
|
"prev_seg_count", prev.numSeg,
|
||||||
|
"prev_total_pay", prev.totalPay,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if canMergeSlots(prev, s) {
|
||||||
|
mergeSlots(prev, s)
|
||||||
|
c.release(s)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
out = append(out, s)
|
||||||
|
}
|
||||||
|
if logged {
|
||||||
|
c.l.Warn("==== end of batch ====")
|
||||||
|
}
|
||||||
|
c.slots = out
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyAddr returns the src or dst address from fk as a netip.Addr for
|
||||||
|
// logging. Only used on the cold gap-log path so the netip allocation
|
||||||
|
// doesn't matter.
|
||||||
|
func flowKeyAddr(fk flowKey, dst bool) netip.Addr {
|
||||||
|
src := fk.src
|
||||||
|
if dst {
|
||||||
|
src = fk.dst
|
||||||
|
}
|
||||||
|
if fk.isV6 {
|
||||||
|
return netip.AddrFrom16(src)
|
||||||
|
}
|
||||||
|
var v4 [4]byte
|
||||||
|
copy(v4[:], src[:4])
|
||||||
|
return netip.AddrFrom4(v4)
|
||||||
|
}
|
||||||
|
|
||||||
|
// sortRun stable-sorts run by (flowKey, seedSeq) so each flow's slots
|
||||||
|
// cluster together in seq order, ready for the merge sweep. Stable so
|
||||||
|
// equal-key slots keep their original relative position (defensive — a
|
||||||
|
// duplicate seedSeq would already mean something's wrong upstream).
|
||||||
|
func (c *TCPCoalescer) sortRun(run []*coalesceSlot) {
|
||||||
|
if len(run) <= 1 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// slices.SortStableFunc with a free, non-capturing comparator avoids the
|
||||||
|
// reflection + closure-escape allocations that sort.SliceStable forces.
|
||||||
|
slices.SortStableFunc(run, compareCoalesceSlots)
|
||||||
|
}
|
||||||
|
|
||||||
|
func compareCoalesceSlots(a, b *coalesceSlot) int {
|
||||||
|
if cmp := flowKeyCompare(a.fk, b.fk); cmp != 0 {
|
||||||
|
return cmp
|
||||||
|
}
|
||||||
|
aSeq, bSeq := slotSeedSeq(a), slotSeedSeq(b)
|
||||||
|
if aSeq == bSeq {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
if tcpSeqLess(aSeq, bSeq) {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
|
// slotSeedSeq returns the TCP seq of the slot's seed (first segment).
|
||||||
|
// nextSeq tracks the seq just past the last appended byte; subtracting
|
||||||
|
// totalPay walks back to the seed. uint32 wraparound is the right TCP
|
||||||
|
// arithmetic so no special-casing is needed.
|
||||||
|
func slotSeedSeq(s *coalesceSlot) uint32 {
|
||||||
|
return s.nextSeq - uint32(s.totalPay)
|
||||||
|
}
|
||||||
|
|
||||||
|
// tcpSeqLess reports whether a precedes b in TCP serial-number arithmetic
|
||||||
|
// (RFC 1323 §2.3). The signed int32 cast turns the modular subtraction
|
||||||
|
// into the right comparison even across the 2^32 wrap.
|
||||||
|
func tcpSeqLess(a, b uint32) bool {
|
||||||
|
return int32(a-b) < 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyCompare orders flowKeys deterministically. The exact ordering
|
||||||
|
// is irrelevant — only that same-flow slots cluster together so the
|
||||||
|
// post-sort sweep can merge contiguous pairs.
|
||||||
|
func flowKeyCompare(a, b flowKey) int {
|
||||||
|
// Cheap scalar fields first so most non-matching keys short-circuit
|
||||||
|
// without ever calling bytes.Compare. sport is the ephemeral port on
|
||||||
|
// egress flows and discriminates fastest. For matching keys (same
|
||||||
|
// flow), array equality on src/dst inlines to word-sized compares,
|
||||||
|
// so we only pay bytes.Compare when the arrays actually differ.
|
||||||
|
if a.sport != b.sport {
|
||||||
|
if a.sport < b.sport {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
if a.dport != b.dport {
|
||||||
|
if a.dport < b.dport {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
if a.dst != b.dst {
|
||||||
|
return bytes.Compare(a.dst[:], b.dst[:])
|
||||||
|
}
|
||||||
|
if a.src != b.src {
|
||||||
|
return bytes.Compare(a.src[:], b.src[:])
|
||||||
|
}
|
||||||
|
if a.isV6 != b.isV6 {
|
||||||
|
if !a.isV6 {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// canMergeSlots reports whether s can fold into prev as one merged TSO
|
||||||
|
// superpacket. Same flow, contiguous TCP byte range, equal gsoSize, and
|
||||||
|
// fits within the kernel TSO limits. The tail-of-prev check rejects any
|
||||||
|
// merge whose first slot ended on a sub-gsoSize segment — kernel TSO
|
||||||
|
// would split the merged skb at gsoSize boundaries from the start, so a
|
||||||
|
// short segment in the middle would corrupt every later segment. PSH and
|
||||||
|
// ECE state must agree across both slots: PSH is a semantic delimiter
|
||||||
|
// (preserving the sender's push boundary) and ECE state must be uniform
|
||||||
|
// across a window (the same rule canAppend enforces for in-flow appends).
|
||||||
|
// The IP-level ECN codepoint must also match: this check calls headersMatch
|
||||||
|
// → ipHeadersMatch, which compares the full DSCP/ECN byte, so two slots with
|
||||||
|
// differing ECN marks stay separate superpackets, each keeping its own mark.
|
||||||
|
//
|
||||||
|
// Note: a slot sealed by reorder (canAppend returned false on seq
|
||||||
|
// mismatch) keeps psh=false, so this restriction does not block the
|
||||||
|
// reorder-fix merge — only legitimate PSH-set seals.
|
||||||
|
func canMergeSlots(prev, s *coalesceSlot) bool {
|
||||||
|
if prev.psh {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.fk != s.fk {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.gsoSize != s.gsoSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.nextSeq != slotSeedSeq(s) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.numSeg+s.numSeg > tcpCoalesceMaxSegs {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.hdrLen+prev.totalPay+s.totalPay > tcpCoalesceBufSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if len(prev.payIovs[len(prev.payIovs)-1]) != prev.gsoSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
prevFlags := prev.hdrBuf[prev.ipHdrLen+13]
|
||||||
|
sFlags := s.hdrBuf[s.ipHdrLen+13]
|
||||||
|
if (prevFlags^sFlags)&tcpFlagEce != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !headersMatch(prev.hdrBuf[:prev.hdrLen], s.hdrBuf[:s.hdrLen], prev.isV6, prev.ipHdrLen) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// mergeSlots folds src into dst in place: payIovs concatenated, counters
|
||||||
|
// and totals updated, PSH OR'd into the seed header so the push signal is
|
||||||
|
// not lost. The seed header's seq, gsoSize, and fk are unchanged. Caller
|
||||||
|
// is responsible for releasing src (it's no longer in c.slots after this call).
|
||||||
|
func mergeSlots(dst, src *coalesceSlot) {
|
||||||
|
dst.payIovs = append(dst.payIovs, src.payIovs...)
|
||||||
|
dst.numSeg += src.numSeg
|
||||||
|
dst.totalPay += src.totalPay
|
||||||
|
dst.nextSeq = src.nextSeq
|
||||||
|
if src.psh {
|
||||||
|
dst.psh = true
|
||||||
|
dst.hdrBuf[dst.ipHdrLen+13] |= tcpFlagPsh
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
|
||||||
|
// already have its checksum field zeroed) and returns the folded/inverted
|
||||||
|
// 16-bit value to store.
|
||||||
|
func ipv4HdrChecksum(hdr []byte) uint16 {
|
||||||
|
var sum uint32
|
||||||
|
for i := 0; i+1 < len(hdr); i += 2 {
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(hdr[i : i+2]))
|
||||||
|
}
|
||||||
|
if len(hdr)%2 == 1 {
|
||||||
|
sum += uint32(hdr[len(hdr)-1]) << 8
|
||||||
|
}
|
||||||
|
for sum>>16 != 0 {
|
||||||
|
sum = (sum & 0xffff) + (sum >> 16)
|
||||||
|
}
|
||||||
|
return ^uint16(sum)
|
||||||
|
}
|
||||||
|
|
||||||
|
// pseudoSumIPv4 / pseudoSumIPv6 build the L4 pseudo-header partial sum
|
||||||
|
// expected by the virtio NEEDS_CSUM kernel path: the 32-bit accumulator
|
||||||
|
// before folding. proto selects the L4 (TCP or UDP); the UDP coalescer
|
||||||
|
// reuses these helpers.
|
||||||
|
func pseudoSumIPv4(src, dst []byte, proto byte, l4Len int) uint32 {
|
||||||
|
var sum uint32
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(src[0:2]))
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(src[2:4]))
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(dst[0:2]))
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(dst[2:4]))
|
||||||
|
sum += uint32(proto)
|
||||||
|
sum += uint32(l4Len)
|
||||||
|
return sum
|
||||||
|
}
|
||||||
|
|
||||||
|
func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
|
||||||
|
var sum uint32
|
||||||
|
for i := 0; i < 16; i += 2 {
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(src[i : i+2]))
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(dst[i : i+2]))
|
||||||
|
}
|
||||||
|
sum += uint32(l4Len >> 16)
|
||||||
|
sum += uint32(l4Len & 0xffff)
|
||||||
|
sum += uint32(proto)
|
||||||
|
return sum
|
||||||
|
}
|
||||||
|
|
||||||
|
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it
|
||||||
|
// unchanged (no one's complement). This is what virtio NEEDS_CSUM wants in
|
||||||
|
// the L4 checksum field — the kernel will add the payload sum and invert.
|
||||||
|
func foldOnceNoInvert(sum uint32) uint16 {
|
||||||
|
for sum>>16 != 0 {
|
||||||
|
sum = (sum & 0xffff) + (sum >> 16)
|
||||||
|
}
|
||||||
|
return uint16(sum)
|
||||||
|
}
|
||||||
@@ -0,0 +1,241 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
)
|
||||||
|
|
||||||
|
// nopTunWriter is a zero-alloc tio.GSOWriter for benchmarks. Discards
|
||||||
|
// everything but satisfies the interface the coalescer detects.
|
||||||
|
type nopTunWriter struct{}
|
||||||
|
|
||||||
|
func (nopTunWriter) Write(p []byte) (int, error) { return len(p), nil }
|
||||||
|
func (nopTunWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, _ tio.GSOProto) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (nopTunWriter) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{TSO: true, USO: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildTCPv4BulkFlow returns a slice of N adjacent ACK-only TCP segments
|
||||||
|
// on a single 5-tuple, each carrying payloadLen bytes. Seq numbers are
|
||||||
|
// contiguous so every packet is coalesceable onto the previous one.
|
||||||
|
func buildTCPv4BulkFlow(n, payloadLen int) [][]byte {
|
||||||
|
pkts := make([][]byte, n)
|
||||||
|
pay := make([]byte, payloadLen)
|
||||||
|
seq := uint32(1000)
|
||||||
|
for i := range n {
|
||||||
|
pkts[i] = buildTCPv4(seq, tcpAck, pay)
|
||||||
|
seq += uint32(payloadLen)
|
||||||
|
}
|
||||||
|
return pkts
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildTCPv4Interleaved returns nFlows * perFlow packets with per-flow
|
||||||
|
// seq continuity but round-robin across flows — worst case for any
|
||||||
|
// "last-slot" cache.
|
||||||
|
func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte {
|
||||||
|
pay := make([]byte, payloadLen)
|
||||||
|
seqs := make([]uint32, nFlows)
|
||||||
|
for i := range seqs {
|
||||||
|
seqs[i] = uint32(1000 + i*1000000)
|
||||||
|
}
|
||||||
|
pkts := make([][]byte, 0, nFlows*perFlow)
|
||||||
|
for range perFlow {
|
||||||
|
for f := range nFlows {
|
||||||
|
sport := uint16(10000 + f)
|
||||||
|
pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay))
|
||||||
|
seqs[f] += uint32(payloadLen)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return pkts
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildICMPv4 returns a minimal non-TCP packet that takes the passthrough
|
||||||
|
// branch in Commit.
|
||||||
|
func buildICMPv4() []byte {
|
||||||
|
pkt := make([]byte, 28)
|
||||||
|
pkt[0] = 0x45
|
||||||
|
binary.BigEndian.PutUint16(pkt[2:4], 28)
|
||||||
|
pkt[9] = 1 // ICMP
|
||||||
|
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||||
|
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||||
|
return pkt
|
||||||
|
}
|
||||||
|
|
||||||
|
// runCommitBench drives Commit over pkts batchSize at a time, flushing
|
||||||
|
// between batches, and reports per-packet cost.
|
||||||
|
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||||
|
b.Helper()
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewTCPCoalescer(nopTunWriter{}, test.NewLogger(), arena.Reserve, arena.Reset)
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.SetBytes(int64(len(pkts[0])))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
pkt := pkts[i%len(pkts)]
|
||||||
|
if err := c.Commit(pkt); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
if (i+1)%batchSize == 0 {
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Drain any trailing partial batch so slot state doesn't leak across runs.
|
||||||
|
_ = c.Flush()
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkCommitSingleFlow is the bulk-TCP steady state: one flow,
|
||||||
|
// contiguous seq, 1200-byte payloads. Every packet past the seed should
|
||||||
|
// append onto the open slot. This is the case we most care about.
|
||||||
|
func BenchmarkCommitSingleFlow(b *testing.B) {
|
||||||
|
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
|
||||||
|
runCommitBench(b, pkts, tcpCoalesceMaxSegs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkCommitInterleaved4 has 4 concurrent bulk flows round-robined.
|
||||||
|
// A single-entry fast-path cache will miss on every packet; an N-way
|
||||||
|
// cache or map lookup carries the weight.
|
||||||
|
func BenchmarkCommitInterleaved4(b *testing.B) {
|
||||||
|
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
||||||
|
runCommitBench(b, pkts, len(pkts))
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkCommitInterleaved16 stresses the map at higher flow counts.
|
||||||
|
func BenchmarkCommitInterleaved16(b *testing.B) {
|
||||||
|
pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200)
|
||||||
|
runCommitBench(b, pkts, len(pkts))
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseTCPBase
|
||||||
|
// bails early and addPassthrough is the only work.
|
||||||
|
func BenchmarkCommitPassthrough(b *testing.B) {
|
||||||
|
pkt := buildICMPv4()
|
||||||
|
pkts := make([][]byte, 64)
|
||||||
|
for i := range pkts {
|
||||||
|
pkts[i] = pkt
|
||||||
|
}
|
||||||
|
runCommitBench(b, pkts, 64)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow.
|
||||||
|
// Each packet takes the "TCP but not admissible" branch which does a
|
||||||
|
// map delete + passthrough. Measures the seal-without-slot cost.
|
||||||
|
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
||||||
|
pay := make([]byte, 0)
|
||||||
|
pkts := make([][]byte, 64)
|
||||||
|
for i := range pkts {
|
||||||
|
pkts[i] = buildTCPv4(uint32(1000+i), tcpSyn|tcpAck, pay)
|
||||||
|
}
|
||||||
|
runCommitBench(b, pkts, 64)
|
||||||
|
}
|
||||||
|
|
||||||
|
// runMultiCommitBench drives MultiCoalescer.Commit. The dispatcher does
|
||||||
|
// the IP/L4 parse once and passes the parsed struct to the lane, so this
|
||||||
|
// is the bench that shows the savings of skipping the lane's re-parse.
|
||||||
|
func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||||
|
b.Helper()
|
||||||
|
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger(), NewArena(0), true, true)
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.SetBytes(int64(len(pkts[0])))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
pkt := pkts[i%len(pkts)]
|
||||||
|
if err := m.Commit(pkt); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
if (i+1)%batchSize == 0 {
|
||||||
|
if err := m.Flush(); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ = m.Flush()
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkMultiCommitSingleFlow is the multi-lane analogue of
|
||||||
|
// BenchmarkCommitSingleFlow — same workload but routed through the
|
||||||
|
// dispatcher. The delta vs the single-lane bench measures dispatcher
|
||||||
|
// overhead.
|
||||||
|
func BenchmarkMultiCommitSingleFlow(b *testing.B) {
|
||||||
|
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
|
||||||
|
runMultiCommitBench(b, pkts, tcpCoalesceMaxSegs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkMultiCommitInterleaved4 mirrors BenchmarkCommitInterleaved4
|
||||||
|
// through the dispatcher.
|
||||||
|
func BenchmarkMultiCommitInterleaved4(b *testing.B) {
|
||||||
|
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
||||||
|
runMultiCommitBench(b, pkts, len(pkts))
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyPair is one comparison input for the flowKeyCompare bench.
|
||||||
|
type flowKeyPair struct{ a, b flowKey }
|
||||||
|
|
||||||
|
// makeFlowKey builds an IPv4 flowKey from compact inputs.
|
||||||
|
func makeFlowKey(srcLow, dstLow uint32, sport, dport uint16) flowKey {
|
||||||
|
var fk flowKey
|
||||||
|
binary.BigEndian.PutUint32(fk.src[12:16], srcLow)
|
||||||
|
binary.BigEndian.PutUint32(fk.dst[12:16], dstLow)
|
||||||
|
fk.sport = sport
|
||||||
|
fk.dport = dport
|
||||||
|
return fk
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyCases are the workload mixes flowKeyCompare sees in practice.
|
||||||
|
// - sameFlow: equal keys; tests the equal-path cost (sort runs hit this
|
||||||
|
// repeatedly when many segments share a flow).
|
||||||
|
// - sportDiffers: same src/dst/dport, different sport — the typical
|
||||||
|
// "sibling flows from one host to one server" pattern.
|
||||||
|
// - dstDiffers: same src/sport/dport, different dst — outbound to many
|
||||||
|
// servers from a fixed local port.
|
||||||
|
// - allDiffer: every field differs; worst case for short-circuiting.
|
||||||
|
func flowKeyCases() map[string][]flowKeyPair {
|
||||||
|
const n = 64
|
||||||
|
cases := map[string][]flowKeyPair{
|
||||||
|
"sameFlow": make([]flowKeyPair, n),
|
||||||
|
"sportDiffers": make([]flowKeyPair, n),
|
||||||
|
"dstDiffers": make([]flowKeyPair, n),
|
||||||
|
"allDiffer": make([]flowKeyPair, n),
|
||||||
|
}
|
||||||
|
for i := range n {
|
||||||
|
base := makeFlowKey(0x0a000001, 0x0a000002, 40000, 443)
|
||||||
|
cases["sameFlow"][i] = flowKeyPair{a: base, b: base}
|
||||||
|
cases["sportDiffers"][i] = flowKeyPair{
|
||||||
|
a: base,
|
||||||
|
b: makeFlowKey(0x0a000001, 0x0a000002, uint16(40001+i), 443),
|
||||||
|
}
|
||||||
|
cases["dstDiffers"][i] = flowKeyPair{
|
||||||
|
a: base,
|
||||||
|
b: makeFlowKey(0x0a000001, uint32(0x0a000002+i+1), 40000, 443),
|
||||||
|
}
|
||||||
|
cases["allDiffer"][i] = flowKeyPair{
|
||||||
|
a: makeFlowKey(uint32(0x0a000001+i), uint32(0x0a000002+i), uint16(40000+i), uint16(80+i)),
|
||||||
|
b: makeFlowKey(uint32(0x0b000001+i), uint32(0x0b000002+i), uint16(50000+i), uint16(443+i)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return cases
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkFlowKeyCompare measures flowKeyCompare across the workloads
|
||||||
|
// the sort step actually sees. Use this to compare reorderings.
|
||||||
|
func BenchmarkFlowKeyCompare(b *testing.B) {
|
||||||
|
for name, pairs := range flowKeyCases() {
|
||||||
|
b.Run(name, func(b *testing.B) {
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
var sink int
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
p := pairs[i&(len(pairs)-1)]
|
||||||
|
sink += flowKeyCompare(p.a, p.b)
|
||||||
|
}
|
||||||
|
runtime.KeepAlive(sink)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user