mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 12:46:58 +02:00
Compare commits
81 Commits
wirebuffer-v2
...
master
| Author | SHA1 | Date | |
|---|---|---|---|
| 6d124d0441 | |||
| 72bf111209 | |||
| 1617897043 | |||
| f8775bb6ca | |||
| 15f0f0d5d0 | |||
| 7902ce674e | |||
| c2fbe215e6 | |||
| 94ac6db4ca | |||
| a60350e34e | |||
| 58f3b6fda7 | |||
| a99699e370 | |||
| 3615a79b8b | |||
| 147c202c27 | |||
| e290a6892f | |||
| 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 |
@@ -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,11 +10,11 @@ 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@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
@@ -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,12 +32,15 @@ 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@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
@@ -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,16 +76,16 @@ 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@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- 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,19 +14,27 @@ 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@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: add hashicorp source
|
- name: add hashicorp source
|
||||||
@@ -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@v7
|
||||||
# which prevents libvirt (used by the other tests) from working after this point.
|
with:
|
||||||
- name: install virtualbox for i386 test
|
go-version: '1.26'
|
||||||
|
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@v7
|
||||||
|
with:
|
||||||
|
go-version: '1.26'
|
||||||
|
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,11 +18,11 @@ 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@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: build
|
- name: build
|
||||||
@@ -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
|
||||||
|
|||||||
+102
-84
@@ -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@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
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
|
||||||
@@ -34,99 +42,109 @@ jobs:
|
|||||||
- name: golangci-lint
|
- name: golangci-lint
|
||||||
uses: golangci/golangci-lint-action@v9
|
uses: golangci/golangci-lint-action@v9
|
||||||
with:
|
with:
|
||||||
version: v2.5
|
version: v2.12
|
||||||
|
|
||||||
- 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@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.26'
|
||||||
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@v7
|
||||||
|
with:
|
||||||
|
go-version: '1.26'
|
||||||
|
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"
|
||||||
|
|||||||
@@ -7,6 +7,88 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
## [Unreleased]
|
## [Unreleased]
|
||||||
|
|
||||||
|
## [1.11.0] - 2026-07-23
|
||||||
|
|
||||||
|
See the [v1.11.0](https://github.com/slackhq/nebula/milestone/25?closed=1) milestone for a complete list of changes.
|
||||||
|
|
||||||
|
### Breaking
|
||||||
|
|
||||||
|
- Logging has switched from logrus to Go's structured `slog`. Log output changes: levels are upper case
|
||||||
|
(`level=INFO`), trace prints as `level=DEBUG-4`, timestamps are always RFC3339Nano and `logging.timestamp_format`
|
||||||
|
is ignored, and some messages were reworded. Review any log parsing before upgrading. This is also an API break
|
||||||
|
for embedders, as constructors now take a `*slog.Logger`. (#1672, #1734, #1621)
|
||||||
|
- `firewall.inbound_action` and `firewall.outbound_action` (used to set reject vs. drop policy) were each being
|
||||||
|
applied to the opposite direction, that is now corrected. This only affects how blocked packets are answered, not
|
||||||
|
which packets the firewall allows or denies. If you set either of these you are getting the behavior of the other
|
||||||
|
one today and likely want to swap them before upgrading. (#1798)
|
||||||
|
- On Windows, Nebula now installs WFP PERMIT filters for the nebula adapter and the listener port by default. WFP
|
||||||
|
sits below Windows Defender Firewall, so any WDF inbound rules you rely on for either will no longer apply. Set
|
||||||
|
`tun.windows_bypass_wdf` and `listen.windows_bypass_wdf` to false to leave WDF in charge. (#1710)
|
||||||
|
- On Windows, the nebula device is now set to the `private` network category instead of whatever Windows decided,
|
||||||
|
which is usually `Public`. This makes the host firewall less restrictive on the overlay. Set
|
||||||
|
`tun.network_category` to `unset` to keep the old behavior. (#1710)
|
||||||
|
- Reject packets for non-TCP now use ICMP code 13, communication administratively prohibited, instead of code 3,
|
||||||
|
port unreachable. Anything keying off the old code needs updating. (#1766, #1768)
|
||||||
|
- The SSH debug server's profiling commands are now confined to `sshd.sandbox_dir`, which defaults to
|
||||||
|
`$TMP/nebula-debug`. Relative paths resolve inside it and absolute paths outside it are rejected, so anything
|
||||||
|
scripting `start-cpu-profile`, `save-heap-profile`, or `save-mutex-profile` with a path elsewhere needs the
|
||||||
|
directory set. The directory is not created for you. (#1622)
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
- Sign the Windows release binaries. (#1718)
|
||||||
|
- Generate IPv6 reject packets, matching the existing IPv4 behavior. (#1766, #1767, #1768)
|
||||||
|
- Accept `-` in `nebula-cert` to read from stdin or write to stdout. (#1714)
|
||||||
|
- Search for both `config.yml` and `config.yaml` in service and command line modes. (#1717)
|
||||||
|
- Add version labels to the Docker/OCI images. (#1772)
|
||||||
|
- Rebind the listener and re-query lighthouses on macOS when the underlay network changes, so devices moving
|
||||||
|
between wifi and wired or between networks recover without waiting for dead tunnel detection. Controlled by
|
||||||
|
`listen.rebind_on_network_change` (default `true`, not reloadable). (#1816)
|
||||||
|
|
||||||
|
### Changed
|
||||||
|
|
||||||
|
- Reload the firewall when the unsafe networks in the certificate change. (#1719)
|
||||||
|
- Reconfigure, start, and stop the stats listener on a config reload instead of requiring a restart. (#1670)
|
||||||
|
- Update a static host's addresses when they change on reload. (#1713)
|
||||||
|
- Don't require a port on ICMP firewall rules. (#1609)
|
||||||
|
- Connection track ICMP traffic. (#1602)
|
||||||
|
- Return `NODATA` instead of `NXDOMAIN` from the DNS server for a name that exists but has no record of the
|
||||||
|
requested type, so clients that query `AAAA` first (busybox/Alpine) fall through to `A`. (#1668)
|
||||||
|
- Record the local host's details in the DNS server. (#1716)
|
||||||
|
- Install Windows unsafe routes as link routes. (#1709)
|
||||||
|
- Reduce relay handshake log spam, and only log a handshake send error at error level when the remote list
|
||||||
|
changes. (#1733, #1765, #1810)
|
||||||
|
- Start, stop, and reload subsystems (DNS, stats, conntrack, ssh, punchy) cleanly without leaking goroutines. (#1640, #1654, #1661, #1667, #1669, #1708, #1806, #1815)
|
||||||
|
- `Control` is now safe to stop and wait on from any lifecycle state, and a new `Control.Wait` blocks until nebula
|
||||||
|
has fully stopped and returns the first fatal reader error. Failed starts release the udp sockets and tun fd
|
||||||
|
instead of leaking them. (#1794)
|
||||||
|
- Trigger an immediate lighthouse update when reconnecting to or adding a lighthouse instead of waiting for the next update tick. (#1645)
|
||||||
|
- Bring the Darwin and OpenBSD tun implementations in line with the other BSDs. (#1703)
|
||||||
|
- Update to build against go v1.26. (#1818)
|
||||||
|
- Various dependency updates. (#1586, #1587, #1604, #1617, #1618, #1627, #1628, #1629, #1652, #1664, #1665, #1697, #1721, #1732, #1742, #1743, #1750, #1763, #1771, #1782, #1800, #1807)
|
||||||
|
|
||||||
|
### Fixed
|
||||||
|
|
||||||
|
- Fix a data race on a host's remote address that could send packets to the wrong address during a roam. (#1773)
|
||||||
|
- Fix tunnels that could permanently escape connection manager monitoring. (#1752)
|
||||||
|
- Fix a crash when reloading the SSH server's trusted keys. (#1787)
|
||||||
|
- Fix hostmap corruption when a host has multiple overlay addresses. Each address now gets its own list instead of
|
||||||
|
a single shared chain, which also fixes two latent bugs on the add and makePrimary paths. (#1788, #1790)
|
||||||
|
- Apply `remote_allow_list` IPv4 rules to 4-in-6 mapped addresses. (#1786)
|
||||||
|
- Don't panic in the DNS server on a short or empty query name. (#1635)
|
||||||
|
- Advance the replay window on relayed packets so a relay drops replayed frames instead of re-forwarding them. (#1751)
|
||||||
|
- Fix a race in relay state handling. (#1753)
|
||||||
|
- Lock replay window updates so concurrent readers can't corrupt it. (#1802)
|
||||||
|
- Reject malformed handshakes more reliably, including invalid ed25519 key lengths. (#1601, #1756)
|
||||||
|
- Properly handle `closetunnel` packets. (#1638)
|
||||||
|
- Fix an IPv6 extension-header length overflow that could make the firewall parse the wrong protocol and ports. (#1789)
|
||||||
|
- Fix relay re-establishment when a handshake arrives over a relay entry that a one-sided teardown left
|
||||||
|
`Disestablished`, which silently dropped every send until dead tunnel detection forced a re-handshake. (#1805)
|
||||||
|
- Don't build new relay state on a tunnel that was just discarded. (#1796)
|
||||||
|
- Don't delete the wrong pending hostinfo in the handshake manager. (#1811)
|
||||||
|
- Don't call the packet reader after a UDP error on Darwin. (#1755)
|
||||||
|
- Open the FreeBSD tun device non blocking. (#1666)
|
||||||
|
|
||||||
## [1.10.3] - 2026-02-06
|
## [1.10.3] - 2026-02-06
|
||||||
|
|
||||||
### Security
|
### Security
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -227,6 +268,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 +280,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 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
|
||||||
|
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++
|
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
|
||||||
|
|||||||
+30
-5
@@ -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,15 +283,17 @@ 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 !isStdio(*cf.outCertPath) {
|
||||||
if _, err := os.Stat(*cf.outCertPath); err == nil {
|
if _, err := os.Stat(*cf.outCertPath); err == nil {
|
||||||
return fmt.Errorf("refusing to overwrite existing CA cert: %s", *cf.outCertPath)
|
return fmt.Errorf("refusing to overwrite existing CA cert: %s", *cf.outCertPath)
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
var c cert.Certificate
|
var c cert.Certificate
|
||||||
var b []byte
|
var b []byte
|
||||||
@@ -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,12 +69,14 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return fmt.Errorf("error while unmarshaling cert: %s", err)
|
return fmt.Errorf("error while unmarshaling cert: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if !qrToStdout {
|
||||||
if *pf.json {
|
if *pf.json {
|
||||||
jsonCerts = append(jsonCerts, c)
|
jsonCerts = append(jsonCerts, c)
|
||||||
} else {
|
} else {
|
||||||
_, _ = out.Write([]byte(c.String()))
|
_, _ = out.Write([]byte(c.String()))
|
||||||
_, _ = out.Write([]byte("\n"))
|
_, _ = out.Write([]byte("\n"))
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if *pf.outQRPath != "" {
|
if *pf.outQRPath != "" {
|
||||||
b, err := c.MarshalPEM()
|
b, err := c.MarshalPEM()
|
||||||
@@ -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
|
||||||
|
|||||||
+38
-16
@@ -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,17 +291,11 @@ 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 *sf.outCertPath == "" {
|
|
||||||
*sf.outCertPath = *sf.name + ".crt"
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
||||||
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
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 !isStdio(*sf.outKeyPath) {
|
||||||
if _, err := os.Stat(*sf.outKeyPath); err == nil {
|
if _, err := os.Stat(*sf.outKeyPath); err == nil {
|
||||||
return fmt.Errorf("refusing to overwrite existing key: %s", *sf.outKeyPath)
|
return fmt.Errorf("refusing to overwrite existing key: %s", *sf.outKeyPath)
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
err = 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,10 +66,13 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
fmt.Println("-config flag must be set")
|
p, err := config.DefaultPath()
|
||||||
flag.Usage()
|
if err != nil {
|
||||||
|
fmt.Println(err)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
*configPath = p
|
||||||
|
}
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
err := c.Load(*configPath)
|
err := c.Load(*configPath)
|
||||||
@@ -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 {
|
||||||
|
|||||||
@@ -0,0 +1,96 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
cert_test "github.com/slackhq/nebula/cert_test"
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestControlStopClosesOnTimer reproduces the dnclient lifecycle: nebula runs as
|
||||||
|
// a library, and on a config update dnclient calls Stop() in-process to tear the
|
||||||
|
// old instance down before starting a new one. This boots a real nebula (real
|
||||||
|
// blocking UDP sockets, tun disabled), lets it run, then Stop()s it on a timer
|
||||||
|
// and asserts it actually closes. If the reader goroutines parked in recvmmsg
|
||||||
|
// don't wake on Close(), Wait() blocks forever and this fails with a goroutine
|
||||||
|
// dump instead of relying on a process signal to unstick them.
|
||||||
|
func TestControlStopClosesOnTimer(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
before := time.Now().Add(-time.Hour)
|
||||||
|
after := time.Now().Add(time.Hour)
|
||||||
|
ca, _, caKey, caPEM := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, before, after, nil, nil, nil)
|
||||||
|
networks := []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}
|
||||||
|
_, _, keyPEM, certPEM := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, "close-on-timer", before, after, networks, nil, nil)
|
||||||
|
|
||||||
|
caPath := filepath.Join(dir, "ca.pem")
|
||||||
|
certPath := filepath.Join(dir, "cert.pem")
|
||||||
|
keyPath := filepath.Join(dir, "key.pem")
|
||||||
|
require.NoError(t, os.WriteFile(caPath, caPEM, 0o600))
|
||||||
|
require.NoError(t, os.WriteFile(certPath, certPEM, 0o600))
|
||||||
|
require.NoError(t, os.WriteFile(keyPath, keyPEM, 0o600))
|
||||||
|
|
||||||
|
// tun disabled so no device/root is needed; routines: 2 so we exercise the
|
||||||
|
// multi-socket (SO_REUSEPORT) teardown, which is where dnclient runs.
|
||||||
|
configBody := fmt.Sprintf(`
|
||||||
|
pki:
|
||||||
|
ca: %s
|
||||||
|
cert: %s
|
||||||
|
key: %s
|
||||||
|
listen:
|
||||||
|
host: 127.0.0.1
|
||||||
|
port: 0
|
||||||
|
tun:
|
||||||
|
disabled: true
|
||||||
|
firewall:
|
||||||
|
outbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
inbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
routines: 2
|
||||||
|
`, caPath, certPath, keyPath)
|
||||||
|
require.NoError(t, os.WriteFile(filepath.Join(dir, "config.yml"), []byte(configBody), 0o600))
|
||||||
|
|
||||||
|
c := config.NewC(l)
|
||||||
|
require.NoError(t, c.Load(dir))
|
||||||
|
|
||||||
|
ctrl, err := nebula.Main(c, false, "close-on-timer", l, nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, ctrl.Start())
|
||||||
|
|
||||||
|
// Run like a live nebula, then close on a timer, exactly as dnclient does.
|
||||||
|
<-time.NewTimer(5 * time.Second).C
|
||||||
|
|
||||||
|
stopped := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
ctrl.Stop() // closes the udp sockets (shutdown(2)) and the tun
|
||||||
|
ctrl.Wait() // blocks until every reader goroutine has returned
|
||||||
|
close(stopped)
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-stopped:
|
||||||
|
t.Log("nebula closed cleanly on timer")
|
||||||
|
case <-time.After(10 * time.Second):
|
||||||
|
buf := make([]byte, 1<<20)
|
||||||
|
n := runtime.Stack(buf, true)
|
||||||
|
t.Fatalf("nebula did NOT close within 10s of Stop(): a blocking reader never woke\n%s", buf[:n])
|
||||||
|
}
|
||||||
|
}
|
||||||
+7
-5
@@ -50,10 +50,13 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
fmt.Println("-config flag must be set")
|
p, err := config.DefaultPath()
|
||||||
flag.Usage()
|
if err != nil {
|
||||||
|
fmt.Println(err)
|
||||||
os.Exit(1)
|
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))
|
||||||
|
}
|
||||||
+9
-47
@@ -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,8 +44,6 @@ type connectionManager struct {
|
|||||||
inactivityTimeout atomic.Int64
|
inactivityTimeout atomic.Int64
|
||||||
dropInactive atomic.Bool
|
dropInactive atomic.Bool
|
||||||
|
|
||||||
metricsTxPunchy metrics.Counter
|
|
||||||
|
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -57,7 +54,6 @@ func newConnectionManagerFromConfig(l *slog.Logger, c *config.C, hm *HostMap, p
|
|||||||
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
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ func newTestLighthouse() *LightHouse {
|
|||||||
lighthouses := []netip.Addr{}
|
lighthouses := []netip.Addr{}
|
||||||
staticList := map[netip.Addr]struct{}{}
|
staticList := map[netip.Addr]struct{}{}
|
||||||
|
|
||||||
|
lh.localAddrsFn = func(*LocalAllowList) []netip.Addr { return nil }
|
||||||
lh.lighthouses.Store(&lighthouses)
|
lh.lighthouses.Store(&lighthouses)
|
||||||
lh.staticList.Store(&staticList)
|
lh.staticList.Store(&staticList)
|
||||||
|
|
||||||
@@ -64,7 +65,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(test.NewLogger())
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
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 +147,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 +234,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 +359,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
|
||||||
|
|||||||
+57
-4
@@ -2,23 +2,27 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"log/slog"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/handshake"
|
"github.com/slackhq/nebula/handshake"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
)
|
)
|
||||||
|
|
||||||
const ReplayWindow = 1024
|
const ReplayWindow = 1024
|
||||||
|
|
||||||
type ConnectionState struct {
|
type ConnectionState struct {
|
||||||
eKey *NebulaCipherState
|
eKey noiseutil.CipherState
|
||||||
dKey *NebulaCipherState
|
dKey noiseutil.CipherState
|
||||||
myCert cert.Certificate
|
myCert cert.Certificate
|
||||||
peerCert *cert.CachedCertificate
|
peerCert *cert.CachedCertificate
|
||||||
initiator bool
|
initiator bool
|
||||||
messageCounter atomic.Uint64
|
messageCounter atomic.Uint64
|
||||||
window *Bits
|
window *Bits
|
||||||
|
decryptLock sync.Mutex
|
||||||
writeLock sync.Mutex
|
writeLock sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -31,8 +35,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)
|
||||||
@@ -53,3 +57,52 @@ func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
|||||||
func (cs *ConnectionState) Curve() cert.Curve {
|
func (cs *ConnectionState) Curve() cert.Curve {
|
||||||
return cs.myCert.Curve()
|
return cs.myCert.Curve()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out []byte, packet []byte, nb []byte) ([]byte, error) {
|
||||||
|
var err error
|
||||||
|
cs.decryptLock.Lock()
|
||||||
|
result := cs.window.Check(l, messageCounter)
|
||||||
|
cs.decryptLock.Unlock()
|
||||||
|
if !result {
|
||||||
|
return nil, ErrAlreadySeen
|
||||||
|
}
|
||||||
|
|
||||||
|
out, err = cs.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
cs.decryptLock.Lock()
|
||||||
|
result = cs.window.Update(l, messageCounter)
|
||||||
|
cs.decryptLock.Unlock()
|
||||||
|
if !result {
|
||||||
|
return nil, ErrAlreadySeen
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// VerifyRelay verifies AEAD protected (but not encrypted) relay frames. packet must be length-checked by the caller.
|
||||||
|
func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) error {
|
||||||
|
cs.decryptLock.Lock()
|
||||||
|
result := cs.window.Check(l, messageCounter)
|
||||||
|
cs.decryptLock.Unlock()
|
||||||
|
if !result {
|
||||||
|
return ErrAlreadySeen
|
||||||
|
}
|
||||||
|
|
||||||
|
signedPayload := packet[:len(packet)-cs.dKey.Overhead()]
|
||||||
|
signatureValue := packet[len(packet)-cs.dKey.Overhead():]
|
||||||
|
_, err := cs.dKey.DecryptDanger(nil, signedPayload, signatureValue, messageCounter, nb)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
cs.decryptLock.Lock()
|
||||||
|
result = cs.window.Update(l, messageCounter)
|
||||||
|
cs.decryptLock.Unlock()
|
||||||
|
if !result {
|
||||||
|
return ErrAlreadySeen
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|||||||
+60
-26
@@ -53,6 +53,7 @@ type Control struct {
|
|||||||
statsStart func()
|
statsStart func()
|
||||||
dnsStart func()
|
dnsStart func()
|
||||||
lighthouseStart func()
|
lighthouseStart func()
|
||||||
|
networkChangeStart func(rebind func())
|
||||||
connectionManagerStart func(context.Context)
|
connectionManagerStart func(context.Context)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -69,29 +70,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.
|
||||||
@@ -104,6 +105,9 @@ func (c *Control) Start() (func() error, error) {
|
|||||||
if c.dnsStart != nil {
|
if c.dnsStart != nil {
|
||||||
go c.dnsStart()
|
go c.dnsStart()
|
||||||
}
|
}
|
||||||
|
if c.networkChangeStart != nil {
|
||||||
|
go c.networkChangeStart(c.RebindUDPServer)
|
||||||
|
}
|
||||||
if c.connectionManagerStart != nil {
|
if c.connectionManagerStart != nil {
|
||||||
go c.connectionManagerStart(c.ctx)
|
go c.connectionManagerStart(c.ctx)
|
||||||
}
|
}
|
||||||
@@ -114,13 +118,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 +133,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 +161,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,9 +193,20 @@ 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.f.outside.Rebind()
|
c.stateLock.Lock()
|
||||||
|
defer c.stateLock.Unlock()
|
||||||
|
|
||||||
|
if c.state != StateStarted {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// A failure here means we are likely still pinned to the interface we came up on, so the rest of this is
|
||||||
|
// unlikely to help. Say so instead of silently carrying on as if we rebound.
|
||||||
|
if err := c.f.outside.Rebind(); err != nil {
|
||||||
|
c.l.Error("Failed to rebind udp socket", "error", err)
|
||||||
|
}
|
||||||
|
|
||||||
// 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
|
||||||
c.f.lightHouse.SendUpdate()
|
c.f.lightHouse.SendUpdate()
|
||||||
@@ -305,7 +339,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 +384,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,292 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/netip"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gaissmai/bart"
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"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(p []byte) (int, error) {
|
||||||
|
<-d.closedCh
|
||||||
|
return 0, 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) SupportsMultiqueue() bool { return false }
|
||||||
|
func (d *fakeDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
|
return nil, errors.New("unsupported")
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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},
|
||||||
|
readers: make([]io.ReadWriteCloser, 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
|
||||||
|
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||||
|
|
||||||
|
// A second Stop is a harmless no-op
|
||||||
|
c.Stop()
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_WaitBlocksUntilStop(t *testing.T) {
|
||||||
|
c, _, _ := newReadyControl(t)
|
||||||
|
|
||||||
|
done := make(chan error, 1)
|
||||||
|
go func() { done <- c.Wait() }()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
t.Fatal("Wait returned before Stop")
|
||||||
|
case <-time.After(50 * time.Millisecond):
|
||||||
|
}
|
||||||
|
|
||||||
|
c.Stop()
|
||||||
|
select {
|
||||||
|
case err := <-done:
|
||||||
|
require.NoError(t, err)
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("Wait did not return after Stop")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type fakeConn struct {
|
||||||
|
closed bool
|
||||||
|
rebinds int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
|
||||||
|
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||||
|
func (c *fakeConn) ListenOut(_ udp.EncReader) error { return nil }
|
||||||
|
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
||||||
|
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
||||||
|
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
||||||
|
func (c *fakeConn) Close() error { c.closed = true; return nil }
|
||||||
|
|
||||||
|
type multiqueueDevice struct {
|
||||||
|
*fakeDevice
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *multiqueueDevice) SupportsMultiqueue() bool { return true }
|
||||||
|
|
||||||
|
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},
|
||||||
|
readers: make([]io.ReadWriteCloser, 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
|
||||||
|
require.Error(t, c.Start())
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||||
|
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||||
|
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
|
||||||
|
|
||||||
|
// And Wait must not hang on the construction token
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInterface_CloseIsIdempotent(t *testing.T) {
|
||||||
|
dev := newFakeDevice()
|
||||||
|
f := &Interface{
|
||||||
|
inside: dev,
|
||||||
|
l: test.NewLogger(),
|
||||||
|
}
|
||||||
|
f.wg.Add(1)
|
||||||
|
|
||||||
|
require.NoError(t, f.Close())
|
||||||
|
assert.True(t, dev.closed)
|
||||||
|
|
||||||
|
// A second Close must not double release the wg token or the device
|
||||||
|
require.NoError(t, f.Close())
|
||||||
|
require.NoError(t, f.wait())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_FatalErrorReportsThroughWait(t *testing.T) {
|
||||||
|
c, dev, conn := newReadyControl(t)
|
||||||
|
|
||||||
|
// Mirror what Start wires up, without needing real packet readers
|
||||||
|
c.f.triggerShutdown = c.Stop
|
||||||
|
c.state = StateStarted
|
||||||
|
|
||||||
|
boom := errors.New("boom")
|
||||||
|
c.f.onFatal(boom)
|
||||||
|
|
||||||
|
require.ErrorIs(t, c.Wait(), boom)
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
assert.True(t, dev.closed)
|
||||||
|
assert.True(t, conn.closed)
|
||||||
|
|
||||||
|
// A second fatal error must not fire the shutdown again or replace the first
|
||||||
|
c.f.onFatal(errors.New("later"))
|
||||||
|
require.ErrorIs(t, c.Wait(), boom)
|
||||||
|
|
||||||
|
// Wait stays factual, a Stop after the death does not mask the error
|
||||||
|
c.Stop()
|
||||||
|
require.ErrorIs(t, c.Wait(), boom)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_ConcurrentStopAndStart(t *testing.T) {
|
||||||
|
c, _, _ := newReadyControl(t)
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
wg.Go(func() { c.Stop() })
|
||||||
|
}
|
||||||
|
wg.Go(func() { _ = c.Start() })
|
||||||
|
wg.Go(func() {
|
||||||
|
_ = c.Wait()
|
||||||
|
// A returned Wait must always observe the final state, no matter how
|
||||||
|
// the race resolved
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
})
|
||||||
|
wg.Wait()
|
||||||
|
|
||||||
|
// However the race resolves, the control must end fully stopped with no
|
||||||
|
// panic and Wait must observe the final state
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_StartStopLifecycle(t *testing.T) {
|
||||||
|
c, dev, conn := newReadyControl(t)
|
||||||
|
|
||||||
|
require.NoError(t, c.Start())
|
||||||
|
assert.Equal(t, StateStarted, c.State())
|
||||||
|
require.ErrorIs(t, c.Start(), ErrAlreadyStarted)
|
||||||
|
|
||||||
|
// Stop must unpark the reader blocked in the device and release everything
|
||||||
|
c.Stop()
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||||
|
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||||
|
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
|
||||||
|
|
||||||
|
// The reader drained off a closed device, that is not a fatal error
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_RebindIsGatedByState(t *testing.T) {
|
||||||
|
c, _, conn := newReadyControl(t)
|
||||||
|
|
||||||
|
// A rebind before Start reaches nothing, the interface is not up
|
||||||
|
c.RebindUDPServer()
|
||||||
|
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
|
||||||
|
|
||||||
|
require.NoError(t, c.Start())
|
||||||
|
c.RebindUDPServer()
|
||||||
|
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
|
||||||
|
|
||||||
|
// A rebind racing a completed stop must not touch the closed conn
|
||||||
|
c.Stop()
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
c.RebindUDPServer()
|
||||||
|
assert.Equal(t, 1, conn.rebinds, "rebind after stop must be a no-op")
|
||||||
|
}
|
||||||
+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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+21
-1
@@ -108,7 +108,19 @@ func (c *Control) GetVpnAddrs() []netip.Addr {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) GetUDPAddr() netip.AddrPort {
|
func (c *Control) GetUDPAddr() netip.AddrPort {
|
||||||
return c.f.outside.(*udp.TesterConn).Addr
|
return c.f.outside.(*udp.TesterConn).GetAddr()
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetUDPAddr moves this node to a new underlay address, standing in for a laptop waking up on a different
|
||||||
|
// network. Register the new address with the router as well or nothing will route back.
|
||||||
|
func (c *Control) SetUDPAddr(addr netip.AddrPort) {
|
||||||
|
c.f.outside.(*udp.TesterConn).SetAddr(addr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetLocalAddrsFn replaces underlay address discovery so a test can advertise its simulated address instead of
|
||||||
|
// whatever this machine's NICs happen to be. Call it before Start, SendUpdate reads it from the update worker.
|
||||||
|
func (c *Control) SetLocalAddrsFn(fn func(*LocalAllowList) []netip.Addr) {
|
||||||
|
c.f.lightHouse.localAddrsFn = fn
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
|
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
|
||||||
@@ -125,6 +137,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))
|
||||||
|
|||||||
+97
-22
@@ -11,7 +11,6 @@ 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"
|
||||||
)
|
)
|
||||||
@@ -23,7 +22,10 @@ type dnsServer struct {
|
|||||||
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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -94,8 +97,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
|||||||
newAddr := getDnsServerAddr(c)
|
newAddr := getDnsServerAddr(c)
|
||||||
|
|
||||||
d.serverMu.Lock()
|
d.serverMu.Lock()
|
||||||
running := d.server
|
running := d.server != nil
|
||||||
runningStarted := d.started
|
|
||||||
sameAddr := d.addr == newAddr
|
sameAddr := d.addr == newAddr
|
||||||
d.addr = newAddr
|
d.addr = newAddr
|
||||||
d.enabled.Store(enabled)
|
d.enabled.Store(enabled)
|
||||||
@@ -109,29 +111,26 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !enabled {
|
if !enabled {
|
||||||
if running != nil {
|
if running {
|
||||||
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
|
||||||
}
|
}
|
||||||
|
|
||||||
if running == nil {
|
if !running {
|
||||||
// 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 {
|
||||||
}
|
// Stop clears the slot before shutting down, otherwise the Start below can find the dying server and refuse
|
||||||
|
d.Stop()
|
||||||
if sameAddr {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
d.shutdownServer(running, runningStarted, "reload")
|
|
||||||
// Old Start goroutine has now exited; bring up a fresh listener on the
|
|
||||||
// new address.
|
|
||||||
go d.Start()
|
go d.Start()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Refresh the self entry every enabled reload so cert renewals that change our name or VPN addresses are picked up.
|
||||||
|
d.seedSelf()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -162,7 +161,9 @@ func (d *dnsServer) Start() {
|
|||||||
|
|
||||||
started := make(chan struct{})
|
started := make(chan struct{})
|
||||||
d.serverMu.Lock()
|
d.serverMu.Lock()
|
||||||
if d.ctx.Err() != nil {
|
// Re-check enabled under the lock, a disable that raced our check above snapshots the slot under it too.
|
||||||
|
// Two reloads in quick succession can both spawn a Start, the loser would orphan the live listener past Stop
|
||||||
|
if d.ctx.Err() != nil || d.server != nil || !d.enabled.Load() {
|
||||||
d.serverMu.Unlock()
|
d.serverMu.Unlock()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -200,6 +201,14 @@ func (d *dnsServer) Start() {
|
|||||||
close(started)
|
close(started)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Release our slot, unless a reload already replaced us, so a dead listener can't block a future Start
|
||||||
|
d.serverMu.Lock()
|
||||||
|
if d.server == server {
|
||||||
|
d.server = nil
|
||||||
|
d.started = nil
|
||||||
|
}
|
||||||
|
d.serverMu.Unlock()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
d.l.Warn("Failed to run the DNS responder", "error", err)
|
d.l.Warn("Failed to run the DNS responder", "error", err)
|
||||||
}
|
}
|
||||||
@@ -249,6 +258,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 +289,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 +380,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) {
|
||||||
|
|||||||
+295
-4
@@ -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"
|
||||||
@@ -191,14 +194,51 @@ func TestDnsServer_reload_initial_serveDnsWithoutLighthouse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDnsServer_reload_sameAddr_noOp(t *testing.T) {
|
func TestDnsServer_reload_sameAddr_noOp(t *testing.T) {
|
||||||
|
port := freeUDPPort(t)
|
||||||
ds, c := newTestDnsServer(t)
|
ds, c := newTestDnsServer(t)
|
||||||
setDnsConfig(c, "127.0.0.1", "0", true, true)
|
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||||
|
|
||||||
require.NoError(t, ds.reload(c, true))
|
require.NoError(t, ds.reload(c, true))
|
||||||
// No server running yet, no addr change. Reload should not spawn anything.
|
|
||||||
|
go ds.Start()
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
before := ds.server
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
require.NotNil(t, before)
|
||||||
|
|
||||||
|
// Same address, so the running listener must be left alone rather than rebuilt under live queries
|
||||||
require.NoError(t, ds.reload(c, false))
|
require.NoError(t, ds.reload(c, false))
|
||||||
assert.True(t, ds.enabled.Load())
|
assert.True(t, ds.enabled.Load())
|
||||||
assert.Nil(t, ds.server)
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
after := ds.server
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
assert.Same(t, before, after, "a same-address reload must not restart the listener")
|
||||||
|
|
||||||
|
ds.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// The branch the old sameAddr test was accidentally hitting: enabled with nothing running means reload starts it.
|
||||||
|
func TestDnsServer_reload_whenNotRunning_starts(t *testing.T) {
|
||||||
|
port := freeUDPPort(t)
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||||
|
|
||||||
|
// initial only records config, it never starts anything
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
assert.Nil(t, ds.server, "the initial reload must not start a listener")
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
require.NoError(t, ds.reload(c, false))
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
assert.NotNil(t, ds.server, "a reload with nothing running should bring DNS up")
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
ds.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDnsServer_StartStop_lifecycle(t *testing.T) {
|
func TestDnsServer_StartStop_lifecycle(t *testing.T) {
|
||||||
@@ -276,6 +316,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)
|
||||||
@@ -338,3 +464,168 @@ func waitFor(t *testing.T, cond func() bool) {
|
|||||||
}
|
}
|
||||||
t.Fatal("timed out waiting for condition")
|
t.Fatal("timed out waiting for condition")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Two reloads in quick succession, or a HUP before Control.Start, can race two Starts at the same listener.
|
||||||
|
func TestDnsServer_Start_isIdempotent(t *testing.T) {
|
||||||
|
port := freeUDPPort(t)
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
|
||||||
|
go ds.Start()
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
first := ds.server
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
require.NotNil(t, first)
|
||||||
|
|
||||||
|
// If the second Start replaces the tracked server, Stop kills the wrong one and the port leaks
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
ds.Start()
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(time.Second * 5):
|
||||||
|
t.Fatal("second Start never returned")
|
||||||
|
}
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
second := ds.server
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
assert.Same(t, first, second, "a second Start must not replace the running server")
|
||||||
|
|
||||||
|
// The real proof, after Stop the port must actually be free
|
||||||
|
ds.Stop()
|
||||||
|
waitFor(t, func() bool {
|
||||||
|
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
_ = pc.Close()
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// An address change must actually end up listening on the new port. Start's guard refuses when a server is already
|
||||||
|
// installed, so reload has to clear the slot before shutting the old one down.
|
||||||
|
func TestDnsServer_reload_addrChange_restarts(t *testing.T) {
|
||||||
|
first := freeUDPPort(t)
|
||||||
|
second := freeUDPPort(t)
|
||||||
|
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
setDnsConfig(c, "127.0.0.1", first, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
|
||||||
|
go ds.Start()
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
// Cycle a few times, the failure this guards against depends on which goroutine wins serverMu
|
||||||
|
for i := range 8 {
|
||||||
|
want := second
|
||||||
|
if i%2 == 1 {
|
||||||
|
want = first
|
||||||
|
}
|
||||||
|
setDnsConfig(c, "127.0.0.1", want, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, false))
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
srv := ds.server
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
require.NotNil(t, srv, "reload left DNS down instead of restarting it")
|
||||||
|
require.Equal(t, "127.0.0.1:"+want, srv.Addr, "reload should be serving the new address")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Land back on second so the port assertions below are meaningful
|
||||||
|
setDnsConfig(c, "127.0.0.1", second, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, false))
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
// The old port must be released and the new one actually held
|
||||||
|
waitFor(t, func() bool {
|
||||||
|
pc, err := net.ListenPacket("udp", "127.0.0.1:"+first)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
_ = pc.Close()
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
_, err := net.ListenPacket("udp", "127.0.0.1:"+second)
|
||||||
|
require.Error(t, err, "the new address should be bound by the DNS responder")
|
||||||
|
|
||||||
|
ds.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// A listener that dies on its own must release the slot, or a later same-addr reload sees it as running and no-ops.
|
||||||
|
func TestDnsServer_Start_bindFailure_releasesSlot(t *testing.T) {
|
||||||
|
port := freeUDPPort(t)
|
||||||
|
blocker, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
|
||||||
|
ds.Start() // returns once the bind fails
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
assert.Nil(t, ds.server, "a listener that failed to bind must not stay parked in the slot")
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
// With the slot released, a reload can retry once the port frees up
|
||||||
|
require.NoError(t, blocker.Close())
|
||||||
|
require.NoError(t, ds.reload(c, false))
|
||||||
|
waitForBind(t, ds)
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
assert.NotNil(t, ds.server, "a same-addr reload should retry after a failed bind")
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
ds.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// A disable that lands while Start is between its unlocked check and the guard must not leave a listener behind.
|
||||||
|
func TestDnsServer_Start_refusesWhenDisabledUnderLock(t *testing.T) {
|
||||||
|
port := freeUDPPort(t)
|
||||||
|
ds, c := newTestDnsServer(t)
|
||||||
|
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||||
|
require.NoError(t, ds.reload(c, true))
|
||||||
|
require.True(t, ds.enabled.Load())
|
||||||
|
|
||||||
|
// Holding serverMu parks Start on the lock, the only way to land the disable in that window on purpose
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
ds.Start()
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
t.Fatal("Start returned early, the test never exercised the window")
|
||||||
|
case <-time.After(time.Millisecond * 100):
|
||||||
|
}
|
||||||
|
|
||||||
|
// The disable reload's critical section. It sees nothing running, so it never calls Stop.
|
||||||
|
ds.enabled.Store(false)
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(time.Second * 5):
|
||||||
|
t.Fatal("Start never returned")
|
||||||
|
}
|
||||||
|
|
||||||
|
ds.serverMu.Lock()
|
||||||
|
assert.Nil(t, ds.server, "Start must not install a listener a disable already cancelled")
|
||||||
|
ds.serverMu.Unlock()
|
||||||
|
|
||||||
|
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||||
|
require.NoError(t, err, "an orphaned listener is still holding the port")
|
||||||
|
_ = pc.Close()
|
||||||
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+173
-30
@@ -405,7 +405,7 @@ func TestStage1Race(t *testing.T) {
|
|||||||
|
|
||||||
r.Log("Spin until connection manager tears down a tunnel")
|
r.Log("Spin until connection manager tears down a tunnel")
|
||||||
|
|
||||||
for 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)
|
||||||
@@ -453,9 +453,11 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
|
|
||||||
r.Log("Nuke my hostmap")
|
r.Log("Nuke my hostmap")
|
||||||
myHostmap := myControl.GetHostmap()
|
myHostmap := myControl.GetHostmap()
|
||||||
|
myHostmap.Lock()
|
||||||
myHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
myHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
||||||
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||||
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||||
|
myHostmap.Unlock()
|
||||||
|
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again")))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again")))
|
||||||
p = r.RouteForAllUntilTxTun(theirControl)
|
p = r.RouteForAllUntilTxTun(theirControl)
|
||||||
@@ -465,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)
|
||||||
@@ -504,9 +506,11 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
|
|
||||||
r.Log("Nuke my hostmap")
|
r.Log("Nuke my hostmap")
|
||||||
theirHostmap := theirControl.GetHostmap()
|
theirHostmap := theirControl.GetHostmap()
|
||||||
|
theirHostmap.Lock()
|
||||||
theirHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
theirHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
||||||
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||||
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||||
|
theirHostmap.Unlock()
|
||||||
|
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again")))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again")))
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
@@ -517,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)
|
||||||
@@ -628,10 +632,10 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
r.Log("Close the tunnel")
|
r.Log("Close the tunnel")
|
||||||
relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
|
relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
|
||||||
|
|
||||||
start := 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.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail")))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail")))
|
||||||
|
|
||||||
@@ -721,6 +725,70 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRelayHandshakeOverDisestablishedEntry(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
// If them tears down the tunnel while me keeps Established relay state, me's next
|
||||||
|
// handshake flows through the relay with no fresh CreateRelayRequest and lands on
|
||||||
|
// them's Disestablished terminal relay entry. them must re-establish that entry, or
|
||||||
|
// its first transmit deletes its only relay and the tunnel is born transmit-dead:
|
||||||
|
// them can receive but every send is silently dropped.
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them ", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
|
||||||
|
|
||||||
|
// Teach my how to get to the relay and that their can be reached via the relay
|
||||||
|
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||||
|
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||||
|
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
|
|
||||||
|
// Build a router so we don't have to reason who gets which packet
|
||||||
|
r := router.NewR(t, myControl, relayControl, theirControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
|
||||||
|
// Start the servers
|
||||||
|
myControl.Start()
|
||||||
|
relayControl.Start()
|
||||||
|
theirControl.Start()
|
||||||
|
|
||||||
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
oldIdx := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false).LocalIndex
|
||||||
|
|
||||||
|
t.Log("Close the tunnel on them only, marking their relay entry Disestablished")
|
||||||
|
theirControl.CloseTunnel(myVpnIpNet[0].Addr(), true)
|
||||||
|
|
||||||
|
t.Log("Re-handshake from me, riding the still-Established relay state")
|
||||||
|
myControl.ReHandshake(theirVpnIpNet[0].Addr())
|
||||||
|
for {
|
||||||
|
h := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
|
||||||
|
if h != nil && h.LocalIndex != oldIdx && h.RemoteIndex != 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
r.RouteForAllExitFunc(func(*udp.Packet, *nebula.Control) router.ExitType {
|
||||||
|
return router.RouteAndExit
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
hAtThem := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
||||||
|
require.NotNil(t, hAtThem, "them should have completed the relayed handshake")
|
||||||
|
require.Equal(t, []netip.Addr{relayVpnIpNet[0].Addr()}, hAtThem.CurrentRelaysToMe, "them should know a relay for the new tunnel")
|
||||||
|
|
||||||
|
t.Log("Send from them to me; their only relay entry must survive the transmit")
|
||||||
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
require.Never(t, func() bool {
|
||||||
|
h := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
||||||
|
return h == nil || len(h.CurrentRelaysToMe) == 0
|
||||||
|
}, time.Second, 10*time.Millisecond, "them deleted its only relay entry; the tunnel is permanently transmit-dead")
|
||||||
|
|
||||||
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
|
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
r.RenderHostmaps("Final hostmaps", myControl, relayControl, theirControl)
|
||||||
|
}
|
||||||
|
|
||||||
func TestStage1RaceRelays(t *testing.T) {
|
func TestStage1RaceRelays(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||||
@@ -819,18 +887,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")
|
||||||
@@ -924,24 +992,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")
|
||||||
@@ -1029,24 +1097,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")
|
||||||
@@ -1123,7 +1191,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)
|
||||||
@@ -1223,7 +1291,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)
|
||||||
@@ -1535,3 +1603,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()
|
||||||
|
}
|
||||||
|
|||||||
+3
-7
@@ -18,14 +18,10 @@ import (
|
|||||||
// retry mechanism gives the wg.Wait()-driven goroutines a moment to drain
|
// retry mechanism gives the wg.Wait()-driven goroutines a moment to drain
|
||||||
// before failing the assertion.
|
// before failing the assertion.
|
||||||
//
|
//
|
||||||
// IgnoreCurrent is necessary in the parallelized suite: other tests can
|
// Intentionally NOT t.Parallel()'d: concurrent tests would have their own
|
||||||
// leave goroutines mid-shutdown when this one runs (Stop is async, the
|
// goroutines running and trip the assertion.
|
||||||
// wg.Wait() drain is not blocking on test return). We're checking that
|
|
||||||
// *this* test's setup tears down cleanly, not that the whole suite is
|
|
||||||
// idle at this moment. Intentionally NOT t.Parallel()'d for the same
|
|
||||||
// reason — concurrent test goroutines would always show up.
|
|
||||||
func TestNoGoroutineLeaks(t *testing.T) {
|
func TestNoGoroutineLeaks(t *testing.T) {
|
||||||
defer goleak.VerifyNone(t, goleak.IgnoreCurrent())
|
defer goleak.VerifyNone(t)
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
||||||
|
|||||||
@@ -0,0 +1,225 @@
|
|||||||
|
//go:build e2e_testing
|
||||||
|
// +build e2e_testing
|
||||||
|
|
||||||
|
package e2e
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"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"
|
||||||
|
)
|
||||||
|
|
||||||
|
// reportedAddrs is what the lighthouse would hand a peer asking where vpnAddr is.
|
||||||
|
func reportedAddrs(t *testing.T, lh *nebula.Control, vpnAddr netip.Addr) []netip.AddrPort {
|
||||||
|
t.Helper()
|
||||||
|
cm := lh.QueryLighthouse(vpnAddr)
|
||||||
|
if cm == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var out []netip.AddrPort
|
||||||
|
for _, c := range *cm {
|
||||||
|
out = append(out, c.Reported...)
|
||||||
|
out = append(out, c.Learned...)
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// waitForLighthouseMsg routes until a lighthouse message lands on lh, or gives up. Reports whether one arrived.
|
||||||
|
func waitForLighthouseMsg(t *testing.T, r *router.R, lh *nebula.Control, wait time.Duration) bool {
|
||||||
|
t.Helper()
|
||||||
|
h := &header.H{}
|
||||||
|
return r.RouteForAllExitFuncOrTimeout(wait, func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
|
if c != lh {
|
||||||
|
return router.KeepRouting
|
||||||
|
}
|
||||||
|
// Punches are a single byte and never parse, they are just not what we are after
|
||||||
|
if err := h.Parse(p.Data); err != nil {
|
||||||
|
return router.KeepRouting
|
||||||
|
}
|
||||||
|
if h.Type == header.LightHouse {
|
||||||
|
return router.RouteAndExit
|
||||||
|
}
|
||||||
|
return router.KeepRouting
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// A laptop that changes networks has to tell the lighthouse promptly, otherwise the lighthouse keeps handing peers
|
||||||
|
// the old address and their punches land nowhere. On a long lighthouse interval the only thing that closes that
|
||||||
|
// window is the rebind, which on darwin the network change monitor drives. The e2e build compiles the monitor out,
|
||||||
|
// so we call RebindUDPServer directly, which is the same thing the monitor does.
|
||||||
|
func TestRebindSendsLighthouseUpdate(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||||
|
"lighthouse": m{"am_lighthouse": true},
|
||||||
|
})
|
||||||
|
|
||||||
|
// 600s interval, so nothing scheduled can send an update during this test. A rebind is the only thing that can.
|
||||||
|
myControl, _, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
||||||
|
"lighthouse": m{
|
||||||
|
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||||
|
"interval": 600,
|
||||||
|
},
|
||||||
|
"static_host_map": m{
|
||||||
|
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
r := router.NewR(t, lhControl, myControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
|
||||||
|
lhControl.Start()
|
||||||
|
myControl.Start()
|
||||||
|
|
||||||
|
// Let the startup registration finish, then clear everything it left behind
|
||||||
|
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
|
||||||
|
r.RouteFor(time.Millisecond * 400)
|
||||||
|
|
||||||
|
// Nothing should be talking to the lighthouse on its own now
|
||||||
|
require.False(t, waitForLighthouseMsg(t, r, lhControl, time.Millisecond*200),
|
||||||
|
"nothing should reach the lighthouse before the rebind")
|
||||||
|
|
||||||
|
myControl.RebindUDPServer()
|
||||||
|
|
||||||
|
assert.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5),
|
||||||
|
"a rebind should push an update to the lighthouse rather than waiting out the interval")
|
||||||
|
|
||||||
|
lhControl.Stop()
|
||||||
|
myControl.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// The other half of a rebind: every live tunnel requeries the lighthouse on its next send. That query is what makes
|
||||||
|
// the lighthouse tell the peer to punch toward our new address, which is the part that actually revives a tunnel
|
||||||
|
// whose remote NAT state died while we were on a different network.
|
||||||
|
func TestRebindRequeriesPeersOnNextSend(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||||
|
"lighthouse": m{"am_lighthouse": true},
|
||||||
|
})
|
||||||
|
|
||||||
|
lhCfg := m{
|
||||||
|
"lighthouse": m{
|
||||||
|
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||||
|
"interval": 600,
|
||||||
|
// Without this the peers advertise this machine's real addresses and then try to punch at them,
|
||||||
|
// which the router has no route for.
|
||||||
|
"local_allow_list": m{
|
||||||
|
"10.0.0.0/24": true,
|
||||||
|
"::/0": false,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"static_host_map": m{
|
||||||
|
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", lhCfg)
|
||||||
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", lhCfg)
|
||||||
|
|
||||||
|
r := router.NewR(t, lhControl, myControl, theirControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
|
||||||
|
lhControl.Start()
|
||||||
|
myControl.Start()
|
||||||
|
theirControl.Start()
|
||||||
|
r.RouteFor(time.Millisecond * 500)
|
||||||
|
|
||||||
|
// Point the peers at each other directly, this test is about the rebind and not about lighthouse discovery
|
||||||
|
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
|
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||||
|
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("initial")))
|
||||||
|
r.RouteFor(time.Second)
|
||||||
|
require.NotNil(t, myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false), "expected a tunnel to them")
|
||||||
|
r.RouteFor(time.Millisecond * 300)
|
||||||
|
|
||||||
|
// Assert on what the peer sees rather than on lighthouse traffic. A query for them makes the lighthouse send
|
||||||
|
// them a punch notification, which is the whole point. Our own update to the lighthouse sends them nothing,
|
||||||
|
// so this cannot be satisfied by the update the rebind itself pushes.
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("quiet")))
|
||||||
|
require.False(t, waitForLighthouseMsg(t, r, theirControl, time.Millisecond*300),
|
||||||
|
"an ordinary send should not requery the lighthouse")
|
||||||
|
|
||||||
|
myControl.RebindUDPServer()
|
||||||
|
r.RouteFor(time.Millisecond * 300) // let the update the rebind itself sends pass by
|
||||||
|
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("after rebind")))
|
||||||
|
assert.True(t, waitForLighthouseMsg(t, r, theirControl, time.Second*5),
|
||||||
|
"the first send after a rebind should requery the lighthouse, which then tells the peer to punch at us")
|
||||||
|
|
||||||
|
lhControl.Stop()
|
||||||
|
myControl.Stop()
|
||||||
|
theirControl.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// The scenario this whole thing exists for: a laptop sleeps at the office and wakes up at home on a new address.
|
||||||
|
// Until it tells the lighthouse, the lighthouse keeps handing peers the office address, so their punches land
|
||||||
|
// nowhere and the tunnel stays dead. On a long interval the rebind is the only thing that closes that window.
|
||||||
|
func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||||
|
"lighthouse": m{"am_lighthouse": true},
|
||||||
|
})
|
||||||
|
|
||||||
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
||||||
|
"lighthouse": m{
|
||||||
|
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||||
|
"interval": 600,
|
||||||
|
},
|
||||||
|
"static_host_map": m{
|
||||||
|
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
// Advertise wherever we currently are rather than this machine's real NICs, read fresh each time so a move
|
||||||
|
// is picked up.
|
||||||
|
myControl.SetLocalAddrsFn(func(*nebula.LocalAllowList) []netip.Addr {
|
||||||
|
return []netip.Addr{myControl.GetUDPAddr().Addr()}
|
||||||
|
})
|
||||||
|
|
||||||
|
r := router.NewR(t, lhControl, myControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
|
||||||
|
lhControl.Start()
|
||||||
|
myControl.Start()
|
||||||
|
|
||||||
|
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
|
||||||
|
r.RouteFor(time.Millisecond * 400)
|
||||||
|
|
||||||
|
require.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), myUdpAddr,
|
||||||
|
"the lighthouse should know the address we started on")
|
||||||
|
|
||||||
|
// Wake up somewhere else
|
||||||
|
newAddr := netip.MustParseAddrPort("10.0.0.99:4242")
|
||||||
|
myControl.SetUDPAddr(newAddr)
|
||||||
|
r.AddRoute(newAddr.Addr(), newAddr.Port(), myControl)
|
||||||
|
|
||||||
|
// Nothing has told the lighthouse, and with interval 600 nothing scheduled will
|
||||||
|
r.RouteFor(time.Millisecond * 400)
|
||||||
|
require.NotContains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
|
||||||
|
"the lighthouse should still be handing out the old address before the rebind")
|
||||||
|
|
||||||
|
myControl.RebindUDPServer()
|
||||||
|
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an update after the rebind")
|
||||||
|
r.RouteFor(time.Millisecond * 400)
|
||||||
|
|
||||||
|
assert.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
|
||||||
|
"after the rebind the lighthouse should hand peers our new address")
|
||||||
|
|
||||||
|
lhControl.Stop()
|
||||||
|
myControl.Stop()
|
||||||
|
}
|
||||||
@@ -0,0 +1,136 @@
|
|||||||
|
//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/udp"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestRecoveryTiming measures how long a tunnel takes to come back after the peer stops accepting our traffic,
|
||||||
|
// which is what a laptop waking on a new network looks like from the peer's side: its NAT has no state for where
|
||||||
|
// we are now, so everything we send disappears.
|
||||||
|
//
|
||||||
|
// It is a measurement, not a pass/fail assertion. Recovery is timed to the moment the peer punches back at us,
|
||||||
|
// since that is when its NAT opens and the tunnel is usable again.
|
||||||
|
//
|
||||||
|
// go test -tags e2e_testing -v -run TestRecoveryTiming ./e2e/
|
||||||
|
func TestRecoveryTiming(t *testing.T) {
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
rebind bool
|
||||||
|
}{
|
||||||
|
{"no trigger", false},
|
||||||
|
{"rebind counter", true},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
d, lost := measureRecovery(t, tc.rebind)
|
||||||
|
t.Logf("RESULT %-16s recovered in %-9v (%d packets lost)", tc.name, d.Round(time.Millisecond), lost)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// measureRecovery returns how long until the peer punched back, and how many of our packets died meanwhile. When
|
||||||
|
// rebind is true we call RebindUDPServer once the tunnel goes dark, which is what the darwin network change
|
||||||
|
// monitor does and what iOS has always done. When false, nothing tells nebula anything is wrong.
|
||||||
|
func measureRecovery(t *testing.T, rebind bool) (time.Duration, int) {
|
||||||
|
t.Helper()
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||||
|
"lighthouse": m{"am_lighthouse": true},
|
||||||
|
})
|
||||||
|
|
||||||
|
peerCfg := m{
|
||||||
|
"lighthouse": m{
|
||||||
|
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||||
|
"interval": 600,
|
||||||
|
"local_allow_list": m{
|
||||||
|
"10.0.0.0/24": true,
|
||||||
|
"::/0": false,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"static_host_map": m{
|
||||||
|
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", peerCfg)
|
||||||
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", peerCfg)
|
||||||
|
|
||||||
|
r := router.NewR(t, lhControl, myControl, theirControl)
|
||||||
|
defer r.RenderFlow()
|
||||||
|
defer func() {
|
||||||
|
lhControl.Stop()
|
||||||
|
myControl.Stop()
|
||||||
|
theirControl.Stop()
|
||||||
|
}()
|
||||||
|
|
||||||
|
lhControl.Start()
|
||||||
|
myControl.Start()
|
||||||
|
theirControl.Start()
|
||||||
|
r.RouteFor(time.Millisecond * 500)
|
||||||
|
|
||||||
|
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
|
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||||
|
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("establish")))
|
||||||
|
r.RouteFor(time.Second)
|
||||||
|
if myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false) == nil {
|
||||||
|
t.Fatal("failed to establish the tunnel we are measuring")
|
||||||
|
}
|
||||||
|
r.RouteFor(time.Millisecond * 500)
|
||||||
|
|
||||||
|
// From here the peer's NAT has no state for us, everything we send it disappears
|
||||||
|
start := time.Now()
|
||||||
|
blackholed := 0
|
||||||
|
var recovered time.Duration
|
||||||
|
|
||||||
|
if rebind {
|
||||||
|
myControl.RebindUDPServer()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Keep the tun busy the way someone retrying a stalled connection would
|
||||||
|
stop := make(chan struct{})
|
||||||
|
defer close(stop)
|
||||||
|
go func() {
|
||||||
|
tick := time.NewTicker(time.Millisecond * 200)
|
||||||
|
defer tick.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-stop:
|
||||||
|
return
|
||||||
|
case <-tick.C:
|
||||||
|
myControl.InjectTunPacket(BuildTunUDPPacket(
|
||||||
|
theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("retry")))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
r.RouteForAllExitFuncOrTimeout(time.Second*30, func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
|
if c == theirControl && p.From == myControl.GetUDPAddr() {
|
||||||
|
blackholed++
|
||||||
|
return router.Drop
|
||||||
|
}
|
||||||
|
|
||||||
|
// The peer reaching us directly is the moment its NAT opened, whether that is a punch or a handshake
|
||||||
|
if c == myControl && p.From == theirUdpAddr {
|
||||||
|
recovered = time.Since(start)
|
||||||
|
return router.RouteAndExit
|
||||||
|
}
|
||||||
|
|
||||||
|
return router.KeepRouting
|
||||||
|
})
|
||||||
|
|
||||||
|
if recovered == 0 {
|
||||||
|
t.Fatalf("no recovery within 30s (%d packets blackholed)", blackholed)
|
||||||
|
}
|
||||||
|
return recovered, blackholed
|
||||||
|
}
|
||||||
+140
-21
@@ -114,6 +114,28 @@ type packet struct {
|
|||||||
packet *udp.Packet
|
packet *udp.Packet
|
||||||
tun bool // a packet pulled off a tun device
|
tun bool // a packet pulled off a tun device
|
||||||
rx bool // the packet was received by a udp device
|
rx bool // the packet was received by a udp device
|
||||||
|
|
||||||
|
// h is the nebula header, parsed once when the packet is recorded. parseErr says why there isn't one, which
|
||||||
|
// the flow log reports rather than hiding. Punchy sends a single byte, so an unparseable packet is normal.
|
||||||
|
h header.H
|
||||||
|
parseErr error
|
||||||
|
}
|
||||||
|
|
||||||
|
// fromAddr and toAddr are the addresses this packet actually travelled between. Reading them off the control
|
||||||
|
// instead would misreport the whole history once a test moves a node. Tun packets are synthesized without
|
||||||
|
// addresses, so they fall back to the control.
|
||||||
|
func (p *packet) fromAddr() netip.AddrPort {
|
||||||
|
if p.tun || !p.packet.From.IsValid() {
|
||||||
|
return p.from.GetUDPAddr()
|
||||||
|
}
|
||||||
|
return p.packet.From
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *packet) toAddr() netip.AddrPort {
|
||||||
|
if p.tun || !p.packet.To.IsValid() {
|
||||||
|
return p.to.GetUDPAddr()
|
||||||
|
}
|
||||||
|
return p.packet.To
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *packet) WasReceived() {
|
func (p *packet) WasReceived() {
|
||||||
@@ -131,6 +153,9 @@ const (
|
|||||||
ExitNow ExitType = 1
|
ExitNow ExitType = 1
|
||||||
// RouteAndExit routes this packet and exits immediately afterwards
|
// RouteAndExit routes this packet and exits immediately afterwards
|
||||||
RouteAndExit ExitType = 2
|
RouteAndExit ExitType = 2
|
||||||
|
// Drop discards this packet without delivering it and keeps routing. Use it to simulate a blackhole, such as
|
||||||
|
// a restrictive NAT refusing traffic from an address it has not seen.
|
||||||
|
Drop ExitType = 3
|
||||||
)
|
)
|
||||||
|
|
||||||
type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
|
type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
|
||||||
@@ -141,7 +166,9 @@ type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
|
|||||||
func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
|
||||||
if err := os.MkdirAll("mermaid", 0755); err != nil {
|
// t.Name() contains a slash for subtests, so the flow log can land in a nested directory
|
||||||
|
fn := filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name()))
|
||||||
|
if err := os.MkdirAll(filepath.Dir(fn), 0755); err != nil {
|
||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -152,7 +179,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
|||||||
outNat: make(map[outNatKey]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: fn,
|
||||||
t: t,
|
t: t,
|
||||||
cancelRender: cancel,
|
cancelRender: cancel,
|
||||||
}
|
}
|
||||||
@@ -249,7 +276,7 @@ func (r *R) renderFlow() {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
addr := e.packet.from.GetUDPAddr()
|
addr := e.packet.fromAddr()
|
||||||
if _, ok := participants[addr]; ok {
|
if _, ok := participants[addr]; ok {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -268,7 +295,6 @@ func (r *R) renderFlow() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Print packets
|
// Print packets
|
||||||
h := &header.H{}
|
|
||||||
for _, e := range r.flow {
|
for _, e := range r.flow {
|
||||||
if e.packet == nil {
|
if e.packet == nil {
|
||||||
//fmt.Fprintf(f, " note over %s: %s\n", strings.Join(participantsVals, ", "), e.note)
|
//fmt.Fprintf(f, " note over %s: %s\n", strings.Join(participantsVals, ", "), e.note)
|
||||||
@@ -280,21 +306,22 @@ func (r *R) renderFlow() {
|
|||||||
fmt.Fprintln(f, r.formatUdpPacket(p))
|
fmt.Fprintln(f, r.formatUdpPacket(p))
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
if err := h.Parse(p.packet.Data); err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
line := "--x"
|
line := "--x"
|
||||||
if p.rx {
|
if p.rx {
|
||||||
line = "->>"
|
line = "->>"
|
||||||
}
|
}
|
||||||
|
|
||||||
fmt.Fprintf(f,
|
detail := fmt.Sprintf("%s(%s), index %v, counter: %v",
|
||||||
" %s%s%s: %s(%s), index %v, counter: %v\n",
|
p.h.TypeName(), p.h.SubTypeName(), p.h.RemoteIndex, p.h.MessageCounter)
|
||||||
normalizeName(p.from.GetUDPAddr().String()),
|
if p.parseErr != nil {
|
||||||
|
detail = fmt.Sprintf("unparsed, %v (%d bytes)", p.parseErr, len(p.packet.Data))
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(f, " %s%s%s: %s\n",
|
||||||
|
normalizeName(p.fromAddr().String()),
|
||||||
line,
|
line,
|
||||||
normalizeName(p.to.GetUDPAddr().String()),
|
normalizeName(p.toAddr().String()),
|
||||||
h.TypeName(), h.SubTypeName(), h.RemoteIndex, h.MessageCounter,
|
detail,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -408,21 +435,24 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
|
|||||||
|
|
||||||
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
|
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
|
||||||
|
|
||||||
if len(r.ignoreFlows) > 0 {
|
|
||||||
var h header.H
|
var h header.H
|
||||||
err := h.Parse(p.Data)
|
var parseErr error
|
||||||
if err != nil {
|
if !tun {
|
||||||
panic(err)
|
parseErr = h.Parse(p.Data)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Decide before copying, the copy comes from a freelist and an ignored packet would never be released
|
||||||
for _, i := range r.ignoreFlows {
|
for _, i := range r.ignoreFlows {
|
||||||
if !tun {
|
if tun {
|
||||||
if i.messageType == h.Type && i.subType == h.Subtype {
|
if i.tun.HasValue && i.tun.IsTrue {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
} else if i.tun.HasValue && i.tun.IsTrue {
|
continue
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// A packet we could not parse has no type to match against, so no rule can ignore it
|
||||||
|
if parseErr == nil && i.messageType == h.Type && i.subType == h.Subtype {
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -431,6 +461,8 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
|
|||||||
to: to,
|
to: to,
|
||||||
packet: p.Copy(),
|
packet: p.Copy(),
|
||||||
tun: tun,
|
tun: tun,
|
||||||
|
h: h,
|
||||||
|
parseErr: parseErr,
|
||||||
}
|
}
|
||||||
|
|
||||||
r.flow = append(r.flow, flowEntry{packet: fp})
|
r.flow = append(r.flow, flowEntry{packet: fp})
|
||||||
@@ -660,6 +692,10 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
p.Release()
|
p.Release()
|
||||||
return
|
return
|
||||||
|
|
||||||
|
case Drop:
|
||||||
|
// Record it so the flow log shows the attempt, but never hand it to the receiver
|
||||||
|
r.unlockedInjectFlow(sender, receiver, p, false)
|
||||||
|
|
||||||
case KeepRouting:
|
case KeepRouting:
|
||||||
fp := r.unlockedInjectFlow(sender, receiver, p, false)
|
fp := r.unlockedInjectFlow(sender, receiver, p, false)
|
||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
@@ -690,6 +726,85 @@ func (r *R) RouteUntilAfterMsgType(sender *nebula.Control, msgType header.Messag
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RouteFor routes everything that shows up for the given duration and then returns. Use it to let a test settle
|
||||||
|
// deterministically rather than sleeping and hoping: a single FlushAll races a completing handshake, which queues
|
||||||
|
// more packets right behind it.
|
||||||
|
func (r *R) RouteFor(d time.Duration) {
|
||||||
|
r.RouteForAllExitFuncOrTimeout(d, func(*udp.Packet, *nebula.Control) ExitType {
|
||||||
|
return KeepRouting
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// RouteForAllExitFuncOrTimeout is RouteForAllExitFunc with a deadline, reporting whether whatDo asked to exit
|
||||||
|
// before time ran out. The unbounded version blocks forever on a quiet network, so this is what a test needs to
|
||||||
|
// assert that something does NOT happen, or to route for a fixed settling period.
|
||||||
|
func (r *R) RouteForAllExitFuncOrTimeout(timeout time.Duration, whatDo ExitFunc) bool {
|
||||||
|
sc := make([]reflect.SelectCase, 0, len(r.controls)+1)
|
||||||
|
cm := make([]*nebula.Control, 0, len(r.controls))
|
||||||
|
|
||||||
|
for _, c := range r.controls {
|
||||||
|
sc = append(sc, reflect.SelectCase{
|
||||||
|
Dir: reflect.SelectRecv,
|
||||||
|
Chan: reflect.ValueOf(c.GetUDPTxChan()),
|
||||||
|
Send: reflect.Value{},
|
||||||
|
})
|
||||||
|
cm = append(cm, c)
|
||||||
|
}
|
||||||
|
|
||||||
|
timer := time.NewTimer(timeout)
|
||||||
|
defer timer.Stop()
|
||||||
|
sc = append(sc, reflect.SelectCase{
|
||||||
|
Dir: reflect.SelectRecv,
|
||||||
|
Chan: reflect.ValueOf(timer.C),
|
||||||
|
Send: reflect.Value{},
|
||||||
|
})
|
||||||
|
|
||||||
|
for {
|
||||||
|
x, rx, _ := reflect.Select(sc)
|
||||||
|
if x == len(cm) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
r.Lock()
|
||||||
|
p := rx.Interface().(*udp.Packet)
|
||||||
|
receiver := r.getControl(cm[x].GetUDPAddr(), p.To, p)
|
||||||
|
if receiver == nil {
|
||||||
|
r.Unlock()
|
||||||
|
panic("Can't RouteForAllExitFuncOrTimeout for host: " + p.To.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
e := whatDo(p, receiver)
|
||||||
|
switch e {
|
||||||
|
case ExitNow:
|
||||||
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
|
return true
|
||||||
|
|
||||||
|
case RouteAndExit:
|
||||||
|
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||||
|
receiver.InjectUDPPacket(p)
|
||||||
|
fp.WasReceived()
|
||||||
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
|
return true
|
||||||
|
|
||||||
|
case Drop:
|
||||||
|
// Record it so the flow log shows the attempt, but never hand it to the receiver
|
||||||
|
r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||||
|
|
||||||
|
case KeepRouting:
|
||||||
|
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||||
|
receiver.InjectUDPPacket(p)
|
||||||
|
fp.WasReceived()
|
||||||
|
|
||||||
|
default:
|
||||||
|
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
||||||
|
}
|
||||||
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (r *R) RouteForAllUntilAfterMsgTypeTo(receiver *nebula.Control, msgType header.MessageType, subType header.MessageSubType) {
|
func (r *R) RouteForAllUntilAfterMsgTypeTo(receiver *nebula.Control, msgType header.MessageType, subType header.MessageSubType) {
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, r *nebula.Control) ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, r *nebula.Control) ExitType {
|
||||||
@@ -782,6 +897,10 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
p.Release()
|
p.Release()
|
||||||
return
|
return
|
||||||
|
|
||||||
|
case Drop:
|
||||||
|
// Record it so the flow log shows the attempt, but never hand it to the receiver
|
||||||
|
r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||||
|
|
||||||
case KeepRouting:
|
case KeepRouting:
|
||||||
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
+101
-6
@@ -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
|
||||||
}
|
}
|
||||||
@@ -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")
|
||||||
}
|
}
|
||||||
|
|||||||
+39
-1
@@ -138,6 +138,22 @@ 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
|
||||||
|
|
||||||
|
# On macOS only
|
||||||
|
# macOS scopes the udp socket to the interface it was created on, so moving between networks (wifi to wired,
|
||||||
|
# office to home) leaves Nebula sending out an interface that no longer has a route. When true, Nebula watches
|
||||||
|
# the routing socket and rebinds the listener once the change settles.
|
||||||
|
# iOS does not use this, the host app drives the same rebind itself.
|
||||||
|
# Default true. Not reloadable.
|
||||||
|
#rebind_on_network_change: 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 +179,21 @@ listen:
|
|||||||
|
|
||||||
punchy:
|
punchy:
|
||||||
# Continues to punch inbound/outbound at a regular interval to avoid expiration of firewall nat mappings
|
# Continues to punch inbound/outbound at a regular interval to avoid expiration of firewall nat mappings
|
||||||
|
# This setting is reloadable.
|
||||||
punch: true
|
punch: true
|
||||||
|
|
||||||
# respond means that a node you are trying to reach will connect back out to you if your hole punching fails
|
# respond means that a node you are trying to reach will connect back out to you if your hole punching fails
|
||||||
# this is extremely useful if one node is behind a difficult nat, such as a symmetric NAT
|
# this is extremely useful if one node is behind a difficult nat, such as a symmetric NAT
|
||||||
# Default is false
|
# Default is false
|
||||||
|
# This setting is reloadable.
|
||||||
#respond: true
|
#respond: true
|
||||||
|
|
||||||
# delays a punch response for misbehaving NATs, default is 1 second.
|
# delays a punch response for misbehaving NATs, default is 1 second.
|
||||||
|
# This setting is reloadable.
|
||||||
#delay: 1s
|
#delay: 1s
|
||||||
|
|
||||||
# set the delay before attempting punchy.respond. Default is 5 seconds. respond must be true to take effect.
|
# set the delay before attempting punchy.respond. Default is 5 seconds. respond must be true to take effect.
|
||||||
|
# This setting is reloadable.
|
||||||
#respond_delay: 5s
|
#respond_delay: 5s
|
||||||
|
|
||||||
# Cipher allows you to choose between the available ciphers for your network. Options are chachapoly or aes
|
# Cipher allows you to choose between the available ciphers for your network. Options are chachapoly or aes
|
||||||
@@ -282,6 +302,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
|
||||||
@@ -367,7 +405,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
|
||||||
|
|
||||||
|
|||||||
@@ -8,6 +8,15 @@ Before=sshd.service
|
|||||||
Type=notify
|
Type=notify
|
||||||
NotifyAccess=main
|
NotifyAccess=main
|
||||||
SyslogIdentifier=nebula
|
SyslogIdentifier=nebula
|
||||||
|
|
||||||
|
# Uncomment to run as an unprivileged user with only CAP_NET_ADMIN. Requires a
|
||||||
|
# nebula user that owns the config directory. Add CAP_NET_BIND_SERVICE to both
|
||||||
|
# lines if any listener (lighthouse DNS, listen.port, stats, sshd) binds <1024.
|
||||||
|
#User=nebula
|
||||||
|
#Group=nebula
|
||||||
|
#CapabilityBoundingSet=CAP_NET_ADMIN
|
||||||
|
#AmbientCapabilities=CAP_NET_ADMIN
|
||||||
|
|
||||||
ExecReload=/bin/kill -HUP $MAINPID
|
ExecReload=/bin/kill -HUP $MAINPID
|
||||||
ExecStart=/usr/local/bin/nebula -config /etc/nebula/config.yml
|
ExecStart=/usr/local/bin/nebula -config /etc/nebula/config.yml
|
||||||
Restart=always
|
Restart=always
|
||||||
|
|||||||
+47
-30
@@ -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
|
||||||
@@ -59,7 +59,8 @@ type Firewall struct {
|
|||||||
|
|
||||||
// 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{
|
||||||
@@ -176,7 +176,7 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
|
|||||||
DefaultTimeout: defaultTimeout,
|
DefaultTimeout: defaultTimeout,
|
||||||
routableNetworks: routableNetworks,
|
routableNetworks: routableNetworks,
|
||||||
assignedNetworks: assignedNetworks,
|
assignedNetworks: assignedNetworks,
|
||||||
hasUnsafeNetworks: hasUnsafeNetworks,
|
unsafeNetworks: unsafeNetworks,
|
||||||
l: l,
|
l: l,
|
||||||
|
|
||||||
incomingMetrics: firewallMetrics{
|
incomingMetrics: firewallMetrics{
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
module github.com/slackhq/nebula
|
module github.com/slackhq/nebula
|
||||||
|
|
||||||
go 1.25.0
|
go 1.26.0
|
||||||
|
|
||||||
require (
|
require (
|
||||||
dario.cat/mergo v1.0.2
|
dario.cat/mergo v1.0.2
|
||||||
@@ -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.54.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.57.0
|
||||||
golang.org/x/sync v0.20.0
|
golang.org/x/sync v0.22.0
|
||||||
golang.org/x/sys v0.43.0
|
golang.org/x/sys v0.47.0
|
||||||
golang.org/x/term v0.42.0
|
golang.org/x/term v0.45.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.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
|
||||||
golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
|
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
|
||||||
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.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
|
||||||
golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw=
|
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
|
||||||
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.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
||||||
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
golang.org/x/sync v0.22.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.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||||
golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
golang.org/x/sys v0.47.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.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
|
||||||
golang.org/x/term v0.42.0/go.mod h1:Dq/D+snpsbazcBG5+F9Q1n2rXV8Ma+71xEjTRufARgY=
|
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
|
||||||
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
|
||||||
|
|||||||
+14
-13
@@ -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
|
||||||
@@ -295,7 +295,13 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
|
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
|
||||||
err := hm.outside.WriteTo(stage0, addr)
|
err := hm.outside.WriteTo(stage0, addr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(hm.l).Error("Failed to send handshake message",
|
// These repeat every attempt, so match the success log below and only shout when the remotes changed
|
||||||
|
level := slog.LevelDebug
|
||||||
|
if remotesHaveChanged {
|
||||||
|
level = slog.LevelError
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo.logger(hm.l).Log(context.Background(), level, "Failed to send handshake message",
|
||||||
"udpAddr", addr,
|
"udpAddr", addr,
|
||||||
"initiatorIndex", hostinfo.localIndexId,
|
"initiatorIndex", hostinfo.localIndexId,
|
||||||
"handshake", hsFields,
|
"handshake", hsFields,
|
||||||
@@ -323,7 +329,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 +436,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 +468,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 +490,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,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -534,8 +535,10 @@ func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
|||||||
|
|
||||||
func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
||||||
for _, addr := range hostinfo.vpnAddrs {
|
for _, addr := range hostinfo.vpnAddrs {
|
||||||
|
if cur, ok := hm.vpnIps[addr]; ok && cur.hostinfo == hostinfo {
|
||||||
delete(hm.vpnIps, addr)
|
delete(hm.vpnIps, addr)
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if len(hm.vpnIps) == 0 {
|
if len(hm.vpnIps) == 0 {
|
||||||
hm.vpnIps = map[netip.Addr]*HandshakeHostInfo{}
|
hm.vpnIps = map[netip.Addr]*HandshakeHostInfo{}
|
||||||
@@ -798,7 +801,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 +967,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) {
|
||||||
@@ -1084,7 +1085,7 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos
|
|||||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||||
// We received a valid handshake on this relay, so make sure the relay
|
// We received a valid handshake on this relay, so make sure the relay
|
||||||
// state reflects that, in case it had been marked Disestablished.
|
// state reflects that, in case it had been marked Disestablished.
|
||||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
|
via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established)
|
||||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||||
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
|
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
+157
-105
@@ -60,7 +60,16 @@ type HostMap struct {
|
|||||||
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
|
||||||
|
|
||||||
@@ -282,7 +287,6 @@ type HostInfo struct {
|
|||||||
type ViaSender struct {
|
type ViaSender struct {
|
||||||
UdpAddr netip.AddrPort
|
UdpAddr netip.AddrPort
|
||||||
relayHI *HostInfo // relayHI is the host info object of the relay
|
relayHI *HostInfo // relayHI is the host info object of the relay
|
||||||
remoteIdx uint32 // remoteIdx is the index included in the header of the received packet
|
|
||||||
relay *Relay // relay contains the rest of the relay information, including the PeerIP of the host trying to communicate with us.
|
relay *Relay // relay contains the rest of the relay information, including the PeerIP of the host trying to communicate with us.
|
||||||
IsRelayed bool // IsRelayed is true if the packet was sent through a relay
|
IsRelayed bool // IsRelayed is true if the packet was sent through a relay
|
||||||
}
|
}
|
||||||
@@ -334,6 +338,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 +387,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,86 +447,67 @@ 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 {
|
||||||
|
// Remove this hostinfo from each of its address lists. The lists are independent, so a
|
||||||
|
// sibling is never promoted to an address it does not own and no other list is touched.
|
||||||
|
final := true
|
||||||
|
for _, addr := range hostinfo.vpnAddrs {
|
||||||
|
if list, ok := hm.moreHosts[addr]; ok {
|
||||||
|
list = removeHostInfo(list, hostinfo)
|
||||||
|
hm.unlockedSetHostsForAddr(addr, list)
|
||||||
|
if len(list) > 0 {
|
||||||
|
final = false
|
||||||
|
}
|
||||||
|
} else if existing, ok := hm.Hosts[addr]; ok {
|
||||||
|
if existing == hostinfo {
|
||||||
|
// Common case, the only hostinfo for this address. moreHosts has no entry to clean up.
|
||||||
delete(hm.Hosts, addr)
|
delete(hm.Hosts, addr)
|
||||||
|
} else {
|
||||||
|
// We don't hold this address but another hostinfo does, we still have a tunnel to the peer
|
||||||
|
final = false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Go maps never shrink their buckets, replace fully drained maps so a node that churned
|
||||||
|
// through a large peer count gives the memory back. Same idiom as the index maps below.
|
||||||
if len(hm.Hosts) == 0 {
|
if len(hm.Hosts) == 0 {
|
||||||
hm.Hosts = map[netip.Addr]*HostInfo{}
|
hm.Hosts = map[netip.Addr]*HostInfo{}
|
||||||
}
|
}
|
||||||
|
if len(hm.moreHosts) == 0 {
|
||||||
if hostinfo.next != nil {
|
hm.moreHosts = map[netip.Addr][]*HostInfo{}
|
||||||
// We had more than 1 hostinfo at this vpn addr, promote the next in the list to primary
|
|
||||||
hm.Hosts[addr] = hostinfo.next
|
|
||||||
// It is primary, there is no previous hostinfo now
|
|
||||||
hostinfo.next.prev = nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
} else {
|
|
||||||
// Relink if we were in the middle of multiple hostinfos for this vpn addr
|
|
||||||
if hostinfo.prev != nil {
|
|
||||||
hostinfo.prev.next = hostinfo.next
|
|
||||||
}
|
|
||||||
|
|
||||||
if hostinfo.next != nil {
|
|
||||||
hostinfo.next.prev = hostinfo.prev
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo.next = nil
|
|
||||||
hostinfo.prev = nil
|
|
||||||
|
|
||||||
// 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
|
||||||
hostinfo2, ok := hm.RemoteIndexes[hostinfo.remoteIndexId]
|
hostinfo2, ok := hm.RemoteIndexes[hostinfo.remoteIndexId]
|
||||||
@@ -502,7 +530,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 +539,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 +584,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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
h = h.next
|
|
||||||
|
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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil, nil, errors.New("unable to find host with relay")
|
return nil, nil, errors.New("unable to find host with relay")
|
||||||
@@ -574,20 +615,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 +658,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 +672,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]
|
||||||
|
if !ok {
|
||||||
|
// Common case, the first hostinfo for this address. moreHosts stays empty.
|
||||||
hm.Hosts[vpnAddr] = hostinfo
|
hm.Hosts[vpnAddr] = hostinfo
|
||||||
|
return
|
||||||
if existing != nil && existing != hostinfo {
|
|
||||||
hostinfo.next = existing
|
|
||||||
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 +724,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 +766,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 +789,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) {
|
||||||
|
|||||||
@@ -87,7 +87,7 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
|
|
||||||
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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -103,7 +103,7 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -333,7 +333,7 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
err = f.writers[0].WriteTo(out, via.remote)
|
err = f.writers[0].WriteTo(out, 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)
|
||||||
}
|
}
|
||||||
@@ -344,7 +344,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
|||||||
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 +391,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,12 +403,12 @@ 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,
|
||||||
"udpAddr", remote,
|
"udpAddr", hr,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
+51
-17
@@ -7,6 +7,7 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"slices"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
@@ -14,6 +15,7 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
|
||||||
|
"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"
|
||||||
@@ -213,6 +215,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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -256,17 +261,16 @@ func (f *Interface) activate() error {
|
|||||||
f.readers[i] = reader
|
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() {
|
||||||
@@ -281,13 +285,14 @@ func (f *Interface) run() (func() error, error) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
return func() error {
|
}
|
||||||
|
|
||||||
|
func (f *Interface) wait() error {
|
||||||
f.wg.Wait()
|
f.wg.Wait()
|
||||||
if e := f.fatalErr.Load(); e != nil {
|
if e := f.fatalErr.Load(); e != nil {
|
||||||
return *e
|
return *e
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}, 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
|
||||||
@@ -320,7 +325,10 @@ func (f *Interface) listenOut(i int) {
|
|||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
||||||
})
|
})
|
||||||
|
|
||||||
if err != nil && !f.closed.Load() {
|
// 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)
|
||||||
}
|
}
|
||||||
@@ -339,7 +347,8 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
|||||||
for {
|
for {
|
||||||
n, err := reader.Read(packet)
|
n, err := reader.Read(packet)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -375,13 +384,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
|
||||||
@@ -491,11 +509,7 @@ 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)
|
||||||
|
|
||||||
for {
|
emit := func() {
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
return
|
|
||||||
case <-ticker.C:
|
|
||||||
f.firewall.EmitStats()
|
f.firewall.EmitStats()
|
||||||
f.handshakeManager.EmitStats()
|
f.handshakeManager.EmitStats()
|
||||||
udpStats()
|
udpStats()
|
||||||
@@ -512,6 +526,18 @@ func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
|||||||
certMaxVersion.Update(int64(certState.v1Cert.Version()))
|
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 {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
emit()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -523,9 +549,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 +573,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")
|
||||||
|
}
|
||||||
+279
-8
@@ -4,27 +4,55 @@ 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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
version := int(packet[0] >> 4)
|
||||||
|
switch version {
|
||||||
|
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] {
|
switch packet[9] {
|
||||||
case 6: // tcp
|
case 6: // tcp
|
||||||
return ipv4CreateRejectTCPPacket(packet, out)
|
return ipv4CreateRejectTCPPacket(packet, out)
|
||||||
default:
|
default:
|
||||||
return ipv4CreateRejectICMPPacket(packet, out)
|
return ipv4CreateRejectICMPPacket(packet, out)
|
||||||
}
|
}
|
||||||
|
case ipv6.Version:
|
||||||
|
if len(packet) < ipv6.HeaderLen {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return ipv6CreateRejectPacket(packet, out)
|
||||||
|
default:
|
||||||
|
return nil
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
|
func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
|
||||||
@@ -35,11 +63,16 @@ 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) {
|
||||||
@@ -72,7 +105,7 @@ 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
|
||||||
@@ -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
|
||||||
|
}
|
||||||
|
|||||||
+404
-1
@@ -1,11 +1,13 @@
|
|||||||
package iputil
|
package iputil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/binary"
|
||||||
"net"
|
"net"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
|
"golang.org/x/net/ipv6"
|
||||||
)
|
)
|
||||||
|
|
||||||
func Test_CreateRejectPacket(t *testing.T) {
|
func Test_CreateRejectPacket(t *testing.T) {
|
||||||
@@ -43,7 +45,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 +73,404 @@ func Test_CreateRejectPacket(t *testing.T) {
|
|||||||
assert.NotNil(t, rejectPacket)
|
assert.NotNil(t, rejectPacket)
|
||||||
assert.Len(t, rejectPacket, expectedLen)
|
assert.Len(t, rejectPacket, expectedLen)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacket_NoFragment(t *testing.T) {
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
|
||||||
|
// IPv4: non-zero fragment offset should not generate reject packet
|
||||||
|
h := ipv4.Header{
|
||||||
|
Len: 20,
|
||||||
|
Src: net.IPv4(10, 0, 0, 1),
|
||||||
|
Dst: net.IPv4(10, 0, 0, 2),
|
||||||
|
Protocol: 17, // UDP
|
||||||
|
}
|
||||||
|
b, err := h.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("h.Marshal: %v", err)
|
||||||
|
}
|
||||||
|
b = append(b, make([]byte, 8)...)
|
||||||
|
// Set fragment offset to non-zero (byte 6-7, offset in 8-byte units)
|
||||||
|
b[6] = 0x00
|
||||||
|
b[7] = 0x01
|
||||||
|
assert.Nil(t, CreateRejectPacket(b, out))
|
||||||
|
|
||||||
|
// MF flag with zero offset (first fragment) should still generate reject
|
||||||
|
b[6] = 0x20 // MF flag set
|
||||||
|
b[7] = 0x00
|
||||||
|
assert.NotNil(t, CreateRejectPacket(b, out))
|
||||||
|
|
||||||
|
// Non-fragment should still generate reject packet
|
||||||
|
b[6] = 0x00
|
||||||
|
b[7] = 0x00
|
||||||
|
assert.NotNil(t, CreateRejectPacket(b, out))
|
||||||
|
|
||||||
|
// DF flag only (not a fragment) should still generate reject packet
|
||||||
|
b[6] = 0x40
|
||||||
|
b[7] = 0x00
|
||||||
|
assert.NotNil(t, CreateRejectPacket(b, out))
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_NoFragment(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
|
||||||
|
// IPv6 with Fragment header and non-zero offset should not generate reject
|
||||||
|
fragHeader := []byte{
|
||||||
|
17, // next header: UDP
|
||||||
|
0, // reserved
|
||||||
|
0, 9, // fragment offset=1 (shifted left 3), M=1
|
||||||
|
0, 0, 0, 1, // identification
|
||||||
|
}
|
||||||
|
udpPayload := make([]byte, 8)
|
||||||
|
payload := append(fragHeader, udpPayload...)
|
||||||
|
packet := makeIPv6Packet(src, dst, 44, payload) // next header 44 = Fragment
|
||||||
|
assert.Nil(t, CreateRejectPacket(packet, out))
|
||||||
|
|
||||||
|
// Fragment header with zero offset (first fragment) should still generate reject
|
||||||
|
fragHeader[2] = 0
|
||||||
|
fragHeader[3] = 1 // offset=0, M=1
|
||||||
|
payload = append(fragHeader, udpPayload...)
|
||||||
|
packet = makeIPv6Packet(src, dst, 44, payload)
|
||||||
|
assert.NotNil(t, CreateRejectPacket(packet, out))
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacket_NoICMPError(t *testing.T) {
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
|
||||||
|
// ICMP error types should not generate reject packets
|
||||||
|
icmpErrorTypes := []byte{3, 4, 5, 11, 12}
|
||||||
|
for _, icmpType := range icmpErrorTypes {
|
||||||
|
h := ipv4.Header{
|
||||||
|
Len: 20,
|
||||||
|
Src: net.IPv4(10, 0, 0, 1),
|
||||||
|
Dst: net.IPv4(10, 0, 0, 2),
|
||||||
|
Protocol: 1, // ICMP
|
||||||
|
}
|
||||||
|
|
||||||
|
b, err := h.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("h.Marshal: %v", err)
|
||||||
|
}
|
||||||
|
b = append(b, icmpType, 0, 0, 0, 0, 0, 0, 0)
|
||||||
|
|
||||||
|
rejectPacket := CreateRejectPacket(b, out)
|
||||||
|
assert.Nil(t, rejectPacket, "ICMP type %d should not generate a reject packet", icmpType)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ICMP non-error types should still generate reject packets
|
||||||
|
icmpNonErrorTypes := []byte{0, 8, 13, 14}
|
||||||
|
for _, icmpType := range icmpNonErrorTypes {
|
||||||
|
h := ipv4.Header{
|
||||||
|
Len: 20,
|
||||||
|
Src: net.IPv4(10, 0, 0, 1),
|
||||||
|
Dst: net.IPv4(10, 0, 0, 2),
|
||||||
|
Protocol: 1, // ICMP
|
||||||
|
}
|
||||||
|
|
||||||
|
b, err := h.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("h.Marshal: %v", err)
|
||||||
|
}
|
||||||
|
b = append(b, icmpType, 0, 0, 0, 0, 0, 0, 0)
|
||||||
|
|
||||||
|
rejectPacket := CreateRejectPacket(b, out)
|
||||||
|
assert.NotNil(t, rejectPacket, "ICMP type %d should generate a reject packet", icmpType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte {
|
||||||
|
b := make([]byte, ipv6.HeaderLen+len(payload))
|
||||||
|
b[0] = ipv6.Version << 4
|
||||||
|
binary.BigEndian.PutUint16(b[4:], uint16(len(payload)))
|
||||||
|
b[6] = nextHeader
|
||||||
|
b[7] = 64
|
||||||
|
copy(b[8:24], src.To16())
|
||||||
|
copy(b[24:40], dst.To16())
|
||||||
|
copy(b[ipv6.HeaderLen:], payload)
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_ICMP(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
|
||||||
|
// Small UDP packet: entire original included in body
|
||||||
|
udpPayload := make([]byte, 20)
|
||||||
|
udpPayload[0] = 0x00 // src port high
|
||||||
|
udpPayload[1] = 0x50 // src port low (80)
|
||||||
|
udpPayload[2] = 0x01 // dst port high
|
||||||
|
udpPayload[3] = 0xBB // dst port low (443)
|
||||||
|
packet := makeIPv6Packet(src, dst, 17, udpPayload)
|
||||||
|
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.NotNil(t, rejectPacket)
|
||||||
|
|
||||||
|
// Small packet fits entirely: 40 (ipv6 hdr) + 8 (icmpv6 hdr) + 60 (original)
|
||||||
|
expectedLen := ipv6.HeaderLen + 8 + len(packet)
|
||||||
|
assert.Len(t, rejectPacket, expectedLen)
|
||||||
|
|
||||||
|
// Verify version
|
||||||
|
assert.Equal(t, byte(ipv6.Version<<4), rejectPacket[0]&0xf0)
|
||||||
|
// Verify next header is ICMPv6 (58)
|
||||||
|
assert.Equal(t, byte(58), rejectPacket[6])
|
||||||
|
// Verify src/dst are swapped
|
||||||
|
assert.Equal(t, dst.To16(), net.IP(rejectPacket[8:24]))
|
||||||
|
assert.Equal(t, src.To16(), net.IP(rejectPacket[24:40]))
|
||||||
|
// Verify ICMPv6 type=1 (Dest Unreachable), code=1 (Administratively prohibited)
|
||||||
|
assert.Equal(t, byte(1), rejectPacket[ipv6.HeaderLen])
|
||||||
|
assert.Equal(t, byte(1), rejectPacket[ipv6.HeaderLen+1])
|
||||||
|
// Verify entire original packet is included in body
|
||||||
|
assert.Equal(t, packet, rejectPacket[ipv6.HeaderLen+8:])
|
||||||
|
|
||||||
|
// Large packet: body is truncated to 1000 bytes
|
||||||
|
largePkt := makeIPv6Packet(src, dst, 17, make([]byte, 1200))
|
||||||
|
rejectPacket = CreateRejectPacket(largePkt, out)
|
||||||
|
assert.NotNil(t, rejectPacket)
|
||||||
|
assert.Len(t, rejectPacket, ipv6.HeaderLen+8+1000)
|
||||||
|
assert.Equal(t, largePkt[:1000], rejectPacket[ipv6.HeaderLen+8:])
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_TCP(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
|
||||||
|
// TCP SYN packet (next header 6)
|
||||||
|
tcpPayload := make([]byte, 20)
|
||||||
|
tcpPayload[0] = 0x00 // src port high
|
||||||
|
tcpPayload[1] = 0x50 // src port low (80)
|
||||||
|
tcpPayload[2] = 0x01 // dst port high
|
||||||
|
tcpPayload[3] = 0xBB // dst port low (443)
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[4:], 1000) // seq
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[8:], 0) // ack seq
|
||||||
|
tcpPayload[12] = (20 >> 2) << 4 // data offset
|
||||||
|
tcpPayload[13] = 0b00000010 // SYN flag
|
||||||
|
|
||||||
|
packet := makeIPv6Packet(src, dst, 6, tcpPayload)
|
||||||
|
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.NotNil(t, rejectPacket)
|
||||||
|
|
||||||
|
// Expected: 40 (ipv6 hdr) + 20 (tcp RST)
|
||||||
|
expectedLen := ipv6.HeaderLen + 20
|
||||||
|
assert.Len(t, rejectPacket, expectedLen)
|
||||||
|
|
||||||
|
// Verify version
|
||||||
|
assert.Equal(t, byte(ipv6.Version<<4), rejectPacket[0]&0xf0)
|
||||||
|
// Verify next header is TCP (6)
|
||||||
|
assert.Equal(t, byte(6), rejectPacket[6])
|
||||||
|
// Verify src/dst are swapped
|
||||||
|
assert.Equal(t, dst.To16(), net.IP(rejectPacket[8:24]))
|
||||||
|
assert.Equal(t, src.To16(), net.IP(rejectPacket[24:40]))
|
||||||
|
// Verify ports are swapped
|
||||||
|
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
||||||
|
assert.Equal(t, uint16(443), binary.BigEndian.Uint16(tcpOut[0:2]))
|
||||||
|
assert.Equal(t, uint16(80), binary.BigEndian.Uint16(tcpOut[2:4]))
|
||||||
|
// RST+ACK flags (since input was SYN without ACK)
|
||||||
|
assert.Equal(t, byte(0b00010100), tcpOut[13])
|
||||||
|
// ack_seq = original seq (1000) + SYN (1) + FIN (0) + segment data (0)
|
||||||
|
assert.Equal(t, uint32(1001), binary.BigEndian.Uint32(tcpOut[8:]))
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_TCPWithACK(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
|
||||||
|
// TCP packet with ACK set
|
||||||
|
tcpPayload := make([]byte, 20)
|
||||||
|
tcpPayload[0] = 0x00
|
||||||
|
tcpPayload[1] = 0x50
|
||||||
|
tcpPayload[2] = 0x01
|
||||||
|
tcpPayload[3] = 0xBB
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[4:], 1000) // seq
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[8:], 2000) // ack seq
|
||||||
|
tcpPayload[12] = (20 >> 2) << 4 // data offset
|
||||||
|
tcpPayload[13] = 0b00010000 // ACK flag
|
||||||
|
|
||||||
|
packet := makeIPv6Packet(src, dst, 6, tcpPayload)
|
||||||
|
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.NotNil(t, rejectPacket)
|
||||||
|
|
||||||
|
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
||||||
|
// RST only (no ACK) since input had ACK
|
||||||
|
assert.Equal(t, byte(0b00000100), tcpOut[13])
|
||||||
|
// seq = original ack_seq
|
||||||
|
assert.Equal(t, uint32(2000), binary.BigEndian.Uint32(tcpOut[4:]))
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_NoICMPError(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
|
||||||
|
// ICMPv6 error types (1-4) should not generate reject packets
|
||||||
|
for icmpType := byte(1); icmpType <= 4; icmpType++ {
|
||||||
|
payload := make([]byte, 8)
|
||||||
|
payload[0] = icmpType
|
||||||
|
packet := makeIPv6Packet(src, dst, 58, payload)
|
||||||
|
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.Nil(t, rejectPacket, "ICMPv6 type %d should not generate a reject packet", icmpType)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ICMPv6 non-error types should still generate reject packets
|
||||||
|
nonErrorTypes := []byte{128, 129, 133, 134}
|
||||||
|
for _, icmpType := range nonErrorTypes {
|
||||||
|
payload := make([]byte, 8)
|
||||||
|
payload[0] = icmpType
|
||||||
|
packet := makeIPv6Packet(src, dst, 58, payload)
|
||||||
|
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.NotNil(t, rejectPacket, "ICMPv6 type %d should generate a reject packet", icmpType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_TooShort(t *testing.T) {
|
||||||
|
// Packet too short to be valid IPv6
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
assert.Nil(t, CreateRejectPacket([]byte{0x60}, out))
|
||||||
|
assert.Nil(t, CreateRejectPacket(make([]byte, 39), out))
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_CreateRejectPacketIPv6_ExtensionHeaders(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
|
||||||
|
// IPv6 + Hop-by-Hop extension header + TCP
|
||||||
|
hopByHop := []byte{
|
||||||
|
6, // next header: TCP
|
||||||
|
0, // length (8 bytes total)
|
||||||
|
0, 0, // padding
|
||||||
|
0, 0, 0, 0,
|
||||||
|
}
|
||||||
|
tcpPayload := make([]byte, 20)
|
||||||
|
tcpPayload[0] = 0x00
|
||||||
|
tcpPayload[1] = 0x50
|
||||||
|
tcpPayload[2] = 0x01
|
||||||
|
tcpPayload[3] = 0xBB
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[4:], 1000)
|
||||||
|
binary.BigEndian.PutUint32(tcpPayload[8:], 2000)
|
||||||
|
tcpPayload[12] = (20 >> 2) << 4
|
||||||
|
tcpPayload[13] = 0b00010000 // ACK
|
||||||
|
|
||||||
|
payload := append(hopByHop, tcpPayload...)
|
||||||
|
packet := makeIPv6Packet(src, dst, 0, payload) // next header 0 = Hop-by-Hop
|
||||||
|
|
||||||
|
out := make([]byte, MaxRejectPacketSize)
|
||||||
|
rejectPacket := CreateRejectPacket(packet, out)
|
||||||
|
assert.NotNil(t, rejectPacket)
|
||||||
|
|
||||||
|
// Should produce TCP RST
|
||||||
|
expectedLen := ipv6.HeaderLen + 20
|
||||||
|
assert.Len(t, rejectPacket, expectedLen)
|
||||||
|
assert.Equal(t, byte(6), rejectPacket[6]) // next header is TCP
|
||||||
|
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
||||||
|
assert.Equal(t, byte(0b00000100), tcpOut[13]) // RST only
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateICMPEchoResponse_IPv4(t *testing.T) {
|
||||||
|
// Build a simple IPv4 ICMP Echo Request
|
||||||
|
packet := make([]byte, 28)
|
||||||
|
packet[0] = 0x45 // version 4, IHL 5
|
||||||
|
binary.BigEndian.PutUint16(packet[2:], uint16(28)) // total length
|
||||||
|
packet[8] = 64 // TTL
|
||||||
|
packet[9] = 1 // protocol ICMP
|
||||||
|
copy(packet[12:16], net.IPv4(10, 0, 0, 1).To4()) // src
|
||||||
|
copy(packet[16:20], net.IPv4(10, 0, 0, 2).To4()) // dst
|
||||||
|
packet[20] = 8 // ICMP Echo Request
|
||||||
|
|
||||||
|
out := make([]byte, len(packet))
|
||||||
|
result := CreateICMPEchoResponse(packet, out)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.Equal(t, byte(0x45), result[0])
|
||||||
|
// src/dst swapped
|
||||||
|
assert.Equal(t, net.IPv4(10, 0, 0, 2).To4(), net.IP(result[12:16]))
|
||||||
|
assert.Equal(t, net.IPv4(10, 0, 0, 1).To4(), net.IP(result[16:20]))
|
||||||
|
// ICMP Echo Reply
|
||||||
|
assert.Equal(t, byte(0), result[20])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateICMPEchoResponse_IPv6(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1").To16()
|
||||||
|
dst := net.ParseIP("fd00::2").To16()
|
||||||
|
|
||||||
|
// Build an IPv6 ICMPv6 Echo Request packet
|
||||||
|
// IPv6 header (40 bytes) + ICMPv6 (8 bytes)
|
||||||
|
packet := make([]byte, 48)
|
||||||
|
packet[0] = 0x60 // version 6
|
||||||
|
payloadLen := uint16(8) // ICMPv6 header only
|
||||||
|
binary.BigEndian.PutUint16(packet[4:], payloadLen)
|
||||||
|
packet[6] = 58 // Next Header: ICMPv6
|
||||||
|
packet[7] = 64 // Hop Limit
|
||||||
|
copy(packet[8:24], src) // src address
|
||||||
|
copy(packet[24:40], dst) // dst address
|
||||||
|
|
||||||
|
// ICMPv6 Echo Request
|
||||||
|
icmp := packet[40:]
|
||||||
|
icmp[0] = 128 // type: Echo Request
|
||||||
|
icmp[1] = 0 // code
|
||||||
|
binary.BigEndian.PutUint16(icmp[4:], 1) // identifier
|
||||||
|
binary.BigEndian.PutUint16(icmp[6:], 1) // sequence number
|
||||||
|
|
||||||
|
// Compute correct checksum for the request
|
||||||
|
csum := ipv6PseudoheaderChecksum(src, dst, 58, uint32(payloadLen))
|
||||||
|
binary.BigEndian.PutUint16(icmp[2:], tcpipChecksum(icmp, csum))
|
||||||
|
|
||||||
|
out := make([]byte, len(packet))
|
||||||
|
result := CreateICMPEchoResponse(packet, out)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
|
||||||
|
// Version should still be 6
|
||||||
|
assert.Equal(t, byte(6), result[0]>>4)
|
||||||
|
// src/dst swapped
|
||||||
|
assert.Equal(t, dst, net.IP(result[8:24]))
|
||||||
|
assert.Equal(t, src, net.IP(result[24:40]))
|
||||||
|
// ICMPv6 Echo Reply type
|
||||||
|
assert.Equal(t, byte(129), result[40])
|
||||||
|
|
||||||
|
// Verify checksum is valid (tcpipChecksum returns 0 when data+checksum is correct)
|
||||||
|
respIcmp := result[40:]
|
||||||
|
verifyCsum := ipv6PseudoheaderChecksum(result[8:24], result[24:40], 58, uint32(payloadLen))
|
||||||
|
assert.Equal(t, uint16(0), tcpipChecksum(respIcmp, verifyCsum))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateICMPEchoResponse_IPv6_NotEchoRequest(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1").To16()
|
||||||
|
dst := net.ParseIP("fd00::2").To16()
|
||||||
|
|
||||||
|
packet := make([]byte, 48)
|
||||||
|
packet[0] = 0x60
|
||||||
|
binary.BigEndian.PutUint16(packet[4:], 8)
|
||||||
|
packet[6] = 58
|
||||||
|
packet[7] = 64
|
||||||
|
copy(packet[8:24], src)
|
||||||
|
copy(packet[24:40], dst)
|
||||||
|
|
||||||
|
// ICMPv6 type 1 (Destination Unreachable) - not Echo Request
|
||||||
|
packet[40] = 1
|
||||||
|
|
||||||
|
out := make([]byte, len(packet))
|
||||||
|
result := CreateICMPEchoResponse(packet, out)
|
||||||
|
assert.Nil(t, result)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1").To16()
|
||||||
|
dst := net.ParseIP("fd00::2").To16()
|
||||||
|
|
||||||
|
packet := make([]byte, 48)
|
||||||
|
packet[0] = 0x60
|
||||||
|
binary.BigEndian.PutUint16(packet[4:], 8)
|
||||||
|
packet[6] = 6 // TCP, not ICMPv6
|
||||||
|
packet[7] = 64
|
||||||
|
copy(packet[8:24], src)
|
||||||
|
copy(packet[24:40], dst)
|
||||||
|
|
||||||
|
out := make([]byte, len(packet))
|
||||||
|
result := CreateICMPEchoResponse(packet, out)
|
||||||
|
assert.Nil(t, result)
|
||||||
|
}
|
||||||
|
|||||||
+43
-50
@@ -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,9 +34,12 @@ type LightHouse struct {
|
|||||||
|
|
||||||
myVpnNetworks []netip.Prefix
|
myVpnNetworks []netip.Prefix
|
||||||
myVpnNetworksTable *bart.Lite
|
myVpnNetworksTable *bart.Lite
|
||||||
punchConn udp.Conn
|
|
||||||
punchy *Punchy
|
punchy *Punchy
|
||||||
|
|
||||||
|
// localAddrsFn enumerates the underlay addresses we advertise. It is a field so tests can supply simulated
|
||||||
|
// addresses rather than whatever this machine's NICs happen to be. Set it before Start.
|
||||||
|
localAddrsFn func(*LocalAllowList) []netip.Addr
|
||||||
|
|
||||||
// Local cache of answers from light houses
|
// Local cache of answers from light houses
|
||||||
// map of vpn addr to answers
|
// map of vpn addr to answers
|
||||||
addrMap map[netip.Addr]*RemoteList
|
addrMap map[netip.Addr]*RemoteList
|
||||||
@@ -76,7 +78,6 @@ 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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -105,12 +106,15 @@ 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)),
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
h.localAddrsFn = func(al *LocalAllowList) []netip.Addr {
|
||||||
|
return localAddrs(h.l, al)
|
||||||
|
}
|
||||||
|
|
||||||
lighthouses := make([]netip.Addr, 0)
|
lighthouses := make([]netip.Addr, 0)
|
||||||
h.lighthouses.Store(&lighthouses)
|
h.lighthouses.Store(&lighthouses)
|
||||||
staticList := make(map[netip.Addr]struct{})
|
staticList := make(map[netip.Addr]struct{})
|
||||||
@@ -118,9 +122,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 +280,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 +299,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") {
|
||||||
@@ -908,7 +926,7 @@ func (lh *LightHouse) SendUpdate() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
lal := lh.GetLocalAllowList()
|
lal := lh.GetLocalAllowList()
|
||||||
for _, e := range localAddrs(lh.l, lal) {
|
for _, e := range lh.localAddrsFn(lal) {
|
||||||
if lh.myVpnNetworksTable.Contains(e) {
|
if lh.myVpnNetworksTable.Contains(e) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -1406,58 +1424,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 +1468,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,9 +1508,11 @@ func (d *NebulaMetaDetails) GetRelays() []netip.Addr {
|
|||||||
|
|
||||||
if len(d.RelayVpnAddrs) > 0 {
|
if len(d.RelayVpnAddrs) > 0 {
|
||||||
for _, r := range d.RelayVpnAddrs {
|
for _, r := range d.RelayVpnAddrs {
|
||||||
|
if r != nil {
|
||||||
relays = append(relays, protoAddrToNetAddr(r))
|
relays = append(relays, protoAddrToNetAddr(r))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
}
|
||||||
return relays
|
return relays
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -303,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,
|
||||||
|
|||||||
@@ -55,7 +55,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
}
|
}
|
||||||
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
|
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
|
||||||
|
|
||||||
ssh, err := sshd.NewSSHServer(l.With("subsystem", "sshd"))
|
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
|
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
|
||||||
}
|
}
|
||||||
@@ -130,6 +130,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 +181,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,7 +205,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
handshakeManager := NewHandshakeManager(l, hostMap, lightHouse, udpConns[0], handshakeConfig)
|
handshakeManager := NewHandshakeManager(l, hostMap, lightHouse, udpConns[0], handshakeConfig)
|
||||||
lightHouse.handshakeTrigger = handshakeManager.trigger
|
lightHouse.handshakeTrigger = handshakeManager.trigger
|
||||||
|
|
||||||
ds, err := newDnsServerFromConfig(ctx, l, pki.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)
|
||||||
}
|
}
|
||||||
@@ -240,6 +251,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
|
|
||||||
handshakeManager.f = ifce
|
handshakeManager.f = ifce
|
||||||
go handshakeManager.Run(ctx)
|
go handshakeManager.Run(ctx)
|
||||||
|
|
||||||
|
punchy.Start(ctx, ifce, hostMap, lightHouse)
|
||||||
}
|
}
|
||||||
|
|
||||||
stats, err := newStatsServerFromConfig(ctx, l, c, buildVersion, configTest)
|
stats, err := newStatsServerFromConfig(ctx, l, c, buildVersion, configTest)
|
||||||
@@ -255,6 +268,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
|
|
||||||
attachCommands(l, c, ssh, ifce)
|
attachCommands(l, c, ssh, ifce)
|
||||||
|
|
||||||
|
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
|
||||||
|
|
||||||
return &Control{
|
return &Control{
|
||||||
state: StateReady,
|
state: StateReady,
|
||||||
f: ifce,
|
f: ifce,
|
||||||
@@ -265,6 +280,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
statsStart: stats.Start,
|
statsStart: stats.Start,
|
||||||
dnsStart: ds.Start,
|
dnsStart: ds.Start,
|
||||||
lighthouseStart: lightHouse.StartUpdateWorker,
|
lighthouseStart: lightHouse.StartUpdateWorker,
|
||||||
|
networkChangeStart: networkChanges.Start,
|
||||||
connectionManagerStart: connManager.Start,
|
connectionManagerStart: connManager.Start,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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())
|
||||||
|
}
|
||||||
+144
-203
@@ -20,23 +20,46 @@ const (
|
|||||||
minFwPacketLen = 4
|
minFwPacketLen = 4
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var ErrOutOfWindow = errors.New("out of window packet")
|
||||||
|
|
||||||
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
err := h.Parse(packet)
|
err := h.Parse(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
||||||
|
// TODO: record metrics for rx holepunch/punchy packets?
|
||||||
if len(packet) > 1 {
|
if len(packet) > 1 {
|
||||||
f.l.Info("Error while parsing inbound packet",
|
f.messageMetrics.RxInvalid(1)
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("Error while parsing inbound packet",
|
||||||
"from", via,
|
"from", via,
|
||||||
"error", err,
|
"error", err,
|
||||||
"packet", packet,
|
"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,44 +67,115 @@ 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
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(packet) < header.Len+hostinfo.ConnectionState.dKey.Overhead() {
|
||||||
|
f.messageMetrics.RxInvalid(1)
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("packet too small", "from", via, "length", len(packet))
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// All remaining packets are encrypted
|
||||||
|
if isMessageRelay {
|
||||||
|
// Relay packets are special, this branch should always early-return
|
||||||
|
if err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, nb); err != nil {
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
out, err = hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, out, packet, 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) {
|
switch h.Subtype {
|
||||||
|
case header.MessageNone:
|
||||||
|
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache)
|
||||||
|
default:
|
||||||
|
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
case header.LightHouse:
|
||||||
|
//TODO: assert via is not relayed
|
||||||
|
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
||||||
|
|
||||||
|
case header.Test:
|
||||||
switch h.Subtype {
|
switch h.Subtype {
|
||||||
case header.MessageNone:
|
case header.TestReply:
|
||||||
if !f.decryptToTun(hostinfo, h.MessageCounter, out, packet, fwPacket, nb, q, localCache) {
|
// 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, hostinfo.ConnectionState, hostinfo, out, nb, packet)
|
||||||
|
default:
|
||||||
|
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
case header.MessageRelay:
|
|
||||||
// The entire body is sent as AD, not encrypted.
|
case header.CloseTunnel:
|
||||||
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
hostinfo.logger(f.l).Info("Close tunnel received, tearing down.", "from", via)
|
||||||
// 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
|
f.closeTunnel(hostinfo)
|
||||||
// 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.
|
case header.Control:
|
||||||
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
f.relayManager.HandleControlMsg(hostinfo, out, f)
|
||||||
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
|
default:
|
||||||
if err != nil {
|
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message type seen", "from", via, "header", h)
|
||||||
return
|
|
||||||
}
|
}
|
||||||
// Successfully validated the thing. Get rid of the Relay header.
|
}
|
||||||
signedPayload = signedPayload[header.Len:]
|
|
||||||
|
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) {
|
||||||
|
// Successfully validated the thing. Get rid of the Relay header and the AEAD tag
|
||||||
|
signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||||
// Pull the Roaming parts up here, and return in all call paths.
|
// 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
|
// Track usage of both the HostInfo and the Relay for the received & authenticated packet
|
||||||
@@ -93,8 +187,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
|
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
|
||||||
// its internal mapping. This should never happen.
|
// its internal mapping. This should never happen.
|
||||||
hostinfo.logger(f.l).Error("HostInfo missing remote relay index",
|
hostinfo.logger(f.l).Error("HostInfo missing remote relay index",
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
"relayRemoteIndex", h.RemoteIndex,
|
||||||
"remoteIndex", h.RemoteIndex,
|
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -106,20 +199,18 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
via = ViaSender{
|
via = ViaSender{
|
||||||
UdpAddr: via.UdpAddr,
|
UdpAddr: via.UdpAddr,
|
||||||
relayHI: hostinfo,
|
relayHI: hostinfo,
|
||||||
remoteIdx: relay.RemoteIndex,
|
|
||||||
relay: relay,
|
relay: relay,
|
||||||
IsRelayed: true,
|
IsRelayed: true,
|
||||||
}
|
}
|
||||||
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache)
|
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache)
|
||||||
return
|
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Find the target HostInfo relay object
|
// Find the target HostInfo relay object
|
||||||
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Info("Failed to find target host info by ip",
|
hostinfo.logger(f.l).Info("Failed to find target host info by ip",
|
||||||
"relayTo", relay.PeerAddr,
|
"relayTo", relay.PeerAddr,
|
||||||
|
"relayFrom", hostinfo.vpnAddrs[0],
|
||||||
"error", err,
|
"error", err,
|
||||||
"hostinfo.vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -128,12 +219,18 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
if targetRelay.State == Established {
|
if targetRelay.State == Established {
|
||||||
switch targetRelay.Type {
|
switch targetRelay.Type {
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Forward this packet through the relay tunnel
|
// Forward this packet through the relay tunnel, rebuilding it in place.
|
||||||
// Find the target HostInfo
|
// Encode overwrites the old outer header, and the new AEAD tag lands where the old one was
|
||||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
fwdBuf := packet[:0:len(packet)] // Cap to len(packet) to protect memory from a larger parent buffer
|
||||||
return
|
f.SendVia(targetHI, targetRelay, signedPayload, nb, fwdBuf, true)
|
||||||
case TerminalType:
|
case TerminalType:
|
||||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
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 {
|
} else {
|
||||||
hostinfo.logger(f.l).Info("Unexpected target relay state",
|
hostinfo.logger(f.l).Info("Unexpected target relay state",
|
||||||
@@ -143,116 +240,11 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
case header.LightHouse:
|
|
||||||
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
|
||||||
if !f.handleEncrypted(ci, via, h) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
d, err := f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(f.l).Error("Failed to decrypt lighthouse packet",
|
|
||||||
"error", err,
|
|
||||||
"from", via,
|
|
||||||
"packet", packet,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
//TODO: assert via is not relayed
|
|
||||||
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, d, f)
|
|
||||||
|
|
||||||
// Fallthrough to the bottom to record incoming traffic
|
|
||||||
|
|
||||||
case header.Test:
|
|
||||||
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 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:
|
|
||||||
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)
|
|
||||||
|
|
||||||
f.closeTunnel(hostinfo)
|
|
||||||
return
|
|
||||||
|
|
||||||
case header.Control:
|
|
||||||
if !f.handleEncrypted(ci, via, h) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
d, err := f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(f.l).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)
|
|
||||||
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("Unexpected relay type", "from", via, "relayType", relay.Type)
|
||||||
}
|
}
|
||||||
return
|
|
||||||
}
|
}
|
||||||
|
|
||||||
f.handleHostRoaming(hostinfo, via)
|
|
||||||
|
|
||||||
f.connectionManager.In(hostinfo)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// closeTunnel closes a tunnel locally, it does not send a closeTunnel packet to the remote
|
// closeTunnel closes a tunnel locally, it does not send a closeTunnel packet to the remote
|
||||||
@@ -270,7 +262,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 +275,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 +283,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 +407,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 {
|
||||||
@@ -515,46 +489,14 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []byte, h *header.H, nb []byte) ([]byte, error) {
|
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
var err error
|
err := newPacket(out, true, fwPacket)
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], mc, nb)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if !hostinfo.ConnectionState.window.Update(f.l, mc) {
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hostinfo.logger(f.l).Debug("dropping out of window packet", "header", h)
|
|
||||||
}
|
|
||||||
return nil, errors.New("out of window packet")
|
|
||||||
}
|
|
||||||
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) bool {
|
|
||||||
var err error
|
|
||||||
|
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(f.l).Error("Failed to decrypt packet", "error", err)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
err = newPacket(out, true, fwPacket)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
||||||
"error", err,
|
"error", err,
|
||||||
"packet", out,
|
"packet", out,
|
||||||
)
|
)
|
||||||
return false
|
return
|
||||||
}
|
|
||||||
|
|
||||||
if !hostinfo.ConnectionState.window.Update(f.l, messageCounter) {
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hostinfo.logger(f.l).Debug("dropping out of window packet", "fwPacket", fwPacket)
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
}
|
||||||
|
|
||||||
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
@@ -568,15 +510,13 @@ func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out
|
|||||||
"reason", dropReason,
|
"reason", dropReason,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
return false
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
f.connectionManager.In(hostinfo)
|
|
||||||
_, err = f.readers[q].Write(out)
|
_, err = f.readers[q].Write(out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to write to tun", "error", err)
|
f.l.Error("Failed to write to tun", "error", err)
|
||||||
}
|
}
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) maybeSendRecvError(endpoint netip.AddrPort, index uint32) {
|
func (f *Interface) maybeSendRecvError(endpoint netip.AddrPort, index uint32) {
|
||||||
@@ -620,10 +560,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,358 @@
|
|||||||
|
//go:build !e2e_testing
|
||||||
|
// +build !e2e_testing
|
||||||
|
|
||||||
|
package overlay
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"runtime"
|
||||||
|
"strings"
|
||||||
|
"syscall"
|
||||||
|
"time"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
)
|
||||||
|
|
||||||
|
// networkCategory mirrors NLM_NETWORK_CATEGORY from netlistmgr.h.
|
||||||
|
type networkCategory int32
|
||||||
|
|
||||||
|
const (
|
||||||
|
networkCategoryPublic networkCategory = 0
|
||||||
|
networkCategoryPrivate networkCategory = 1
|
||||||
|
networkCategoryDomainAuthenticated networkCategory = 2
|
||||||
|
)
|
||||||
|
|
||||||
|
func (c networkCategory) String() string {
|
||||||
|
switch c {
|
||||||
|
case networkCategoryPublic:
|
||||||
|
return "public"
|
||||||
|
case networkCategoryPrivate:
|
||||||
|
return "private"
|
||||||
|
case networkCategoryDomainAuthenticated:
|
||||||
|
return "domain"
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("unknown(%d)", c)
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseNetworkCategory accepts the user-supplied tun.network_category. A
|
||||||
|
// second return of false means "leave the category alone".
|
||||||
|
func parseNetworkCategory(s string) (networkCategory, bool, error) {
|
||||||
|
switch strings.ToLower(strings.TrimSpace(s)) {
|
||||||
|
case "", "unset":
|
||||||
|
return 0, false, nil
|
||||||
|
case "public":
|
||||||
|
return networkCategoryPublic, true, nil
|
||||||
|
case "private":
|
||||||
|
return networkCategoryPrivate, true, nil
|
||||||
|
case "domain", "domainauthenticated":
|
||||||
|
return networkCategoryDomainAuthenticated, true, nil
|
||||||
|
}
|
||||||
|
return 0, false, fmt.Errorf("unknown tun.network_category %q (expected public, private, domain, or unset)", s)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CLSID_NetworkListManager {DCB00C01-570F-4A9B-8D69-199FDBA5723B}
|
||||||
|
var clsidNetworkListManager = windows.GUID{
|
||||||
|
Data1: 0xDCB00C01, Data2: 0x570F, Data3: 0x4A9B,
|
||||||
|
Data4: [8]byte{0x8D, 0x69, 0x19, 0x9F, 0xDB, 0xA5, 0x72, 0x3B},
|
||||||
|
}
|
||||||
|
|
||||||
|
// IID_INetworkListManager {DCB00000-570F-4A9B-8D69-199FDBA5723B}
|
||||||
|
var iidINetworkListManager = windows.GUID{
|
||||||
|
Data1: 0xDCB00000, Data2: 0x570F, Data3: 0x4A9B,
|
||||||
|
Data4: [8]byte{0x8D, 0x69, 0x19, 0x9F, 0xDB, 0xA5, 0x72, 0x3B},
|
||||||
|
}
|
||||||
|
|
||||||
|
// x/sys/windows doesn't expose CoCreateInstance, so we bind it ourselves.
|
||||||
|
var procCoCreateInstance = windows.NewLazySystemDLL("ole32.dll").NewProc("CoCreateInstance")
|
||||||
|
|
||||||
|
const clsCtxAll = windows.CLSCTX_INPROC_SERVER | windows.CLSCTX_INPROC_HANDLER |
|
||||||
|
windows.CLSCTX_LOCAL_SERVER | windows.CLSCTX_REMOTE_SERVER
|
||||||
|
|
||||||
|
const (
|
||||||
|
hrSFALSE = 0x00000001
|
||||||
|
hrRPCEChangedMode = 0x80010106
|
||||||
|
)
|
||||||
|
|
||||||
|
type hresult uint32
|
||||||
|
|
||||||
|
func (h hresult) failed() bool { return int32(h) < 0 }
|
||||||
|
func (h hresult) String() string {
|
||||||
|
return fmt.Sprintf("HRESULT 0x%08x", uint32(h))
|
||||||
|
}
|
||||||
|
|
||||||
|
var errAdapterNotFound = errors.New("adapter not present in network connections enumeration")
|
||||||
|
|
||||||
|
// Vtable layouts. Slot order must match the declaration order in netlistmgr.h.
|
||||||
|
// All NLM interfaces here derive from IDispatch, which derives from IUnknown.
|
||||||
|
|
||||||
|
type iUnknownVtbl struct {
|
||||||
|
QueryInterface uintptr
|
||||||
|
AddRef uintptr
|
||||||
|
Release uintptr
|
||||||
|
}
|
||||||
|
|
||||||
|
type iDispatchVtbl struct {
|
||||||
|
iUnknownVtbl
|
||||||
|
GetTypeInfoCount uintptr
|
||||||
|
GetTypeInfo uintptr
|
||||||
|
GetIDsOfNames uintptr
|
||||||
|
Invoke uintptr
|
||||||
|
}
|
||||||
|
|
||||||
|
type iNetworkListManagerVtbl struct {
|
||||||
|
iDispatchVtbl
|
||||||
|
GetNetworks uintptr
|
||||||
|
GetNetwork uintptr
|
||||||
|
GetNetworkConnections uintptr
|
||||||
|
GetNetworkConnection uintptr
|
||||||
|
IsConnectedToInternet uintptr
|
||||||
|
IsConnected uintptr
|
||||||
|
GetConnectivity uintptr
|
||||||
|
}
|
||||||
|
|
||||||
|
type iNetworkListManager struct{ Vtbl *iNetworkListManagerVtbl }
|
||||||
|
|
||||||
|
func (n *iNetworkListManager) Release() {
|
||||||
|
syscall.SyscallN(n.Vtbl.Release, uintptr(unsafe.Pointer(n)))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (n *iNetworkListManager) GetNetworkConnections() (*iEnumNetworkConnections, error) {
|
||||||
|
var enum *iEnumNetworkConnections
|
||||||
|
r1, _, _ := syscall.SyscallN(n.Vtbl.GetNetworkConnections,
|
||||||
|
uintptr(unsafe.Pointer(n)), uintptr(unsafe.Pointer(&enum)),
|
||||||
|
)
|
||||||
|
if hr := hresult(r1); hr.failed() {
|
||||||
|
return nil, fmt.Errorf("INetworkListManager.GetNetworkConnections: %s", hr)
|
||||||
|
}
|
||||||
|
return enum, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type iEnumNetworkConnectionsVtbl struct {
|
||||||
|
iDispatchVtbl
|
||||||
|
NewEnum uintptr
|
||||||
|
Next uintptr
|
||||||
|
Skip uintptr
|
||||||
|
Reset uintptr
|
||||||
|
Clone uintptr
|
||||||
|
}
|
||||||
|
|
||||||
|
type iEnumNetworkConnections struct{ Vtbl *iEnumNetworkConnectionsVtbl }
|
||||||
|
|
||||||
|
func (e *iEnumNetworkConnections) Release() {
|
||||||
|
syscall.SyscallN(e.Vtbl.Release, uintptr(unsafe.Pointer(e)))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Next returns the next connection, or (nil, nil) at the end of the enumeration.
|
||||||
|
func (e *iEnumNetworkConnections) Next() (*iNetworkConnection, error) {
|
||||||
|
var conn *iNetworkConnection
|
||||||
|
var fetched uint32
|
||||||
|
r1, _, _ := syscall.SyscallN(e.Vtbl.Next,
|
||||||
|
uintptr(unsafe.Pointer(e)), 1,
|
||||||
|
uintptr(unsafe.Pointer(&conn)), uintptr(unsafe.Pointer(&fetched)),
|
||||||
|
)
|
||||||
|
if hr := hresult(r1); hr.failed() {
|
||||||
|
return nil, fmt.Errorf("IEnumNetworkConnections.Next: %s", hr)
|
||||||
|
}
|
||||||
|
if fetched == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return conn, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type iNetworkConnectionVtbl struct {
|
||||||
|
iDispatchVtbl
|
||||||
|
GetNetwork uintptr
|
||||||
|
IsConnectedToInternet uintptr
|
||||||
|
IsConnected uintptr
|
||||||
|
GetConnectivity uintptr
|
||||||
|
GetConnectionId uintptr
|
||||||
|
GetAdapterId uintptr
|
||||||
|
GetDomainType uintptr
|
||||||
|
}
|
||||||
|
|
||||||
|
type iNetworkConnection struct{ Vtbl *iNetworkConnectionVtbl }
|
||||||
|
|
||||||
|
func (c *iNetworkConnection) Release() {
|
||||||
|
syscall.SyscallN(c.Vtbl.Release, uintptr(unsafe.Pointer(c)))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *iNetworkConnection) GetAdapterId() (windows.GUID, error) {
|
||||||
|
var g windows.GUID
|
||||||
|
r1, _, _ := syscall.SyscallN(c.Vtbl.GetAdapterId,
|
||||||
|
uintptr(unsafe.Pointer(c)), uintptr(unsafe.Pointer(&g)),
|
||||||
|
)
|
||||||
|
if hr := hresult(r1); hr.failed() {
|
||||||
|
return windows.GUID{}, fmt.Errorf("INetworkConnection.GetAdapterId: %s", hr)
|
||||||
|
}
|
||||||
|
return g, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *iNetworkConnection) GetNetwork() (*iNetwork, error) {
|
||||||
|
var net *iNetwork
|
||||||
|
r1, _, _ := syscall.SyscallN(c.Vtbl.GetNetwork,
|
||||||
|
uintptr(unsafe.Pointer(c)), uintptr(unsafe.Pointer(&net)),
|
||||||
|
)
|
||||||
|
if hr := hresult(r1); hr.failed() {
|
||||||
|
return nil, fmt.Errorf("INetworkConnection.GetNetwork: %s", hr)
|
||||||
|
}
|
||||||
|
return net, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type iNetworkVtbl struct {
|
||||||
|
iDispatchVtbl
|
||||||
|
GetName uintptr
|
||||||
|
SetName uintptr
|
||||||
|
GetDescription uintptr
|
||||||
|
SetDescription uintptr
|
||||||
|
GetNetworkId uintptr
|
||||||
|
GetDomainType uintptr
|
||||||
|
GetNetworkConnections uintptr
|
||||||
|
GetTimeCreatedAndConnected uintptr
|
||||||
|
IsConnectedToInternet uintptr
|
||||||
|
IsConnected uintptr
|
||||||
|
GetConnectivity uintptr
|
||||||
|
GetCategory uintptr
|
||||||
|
SetCategory uintptr
|
||||||
|
}
|
||||||
|
|
||||||
|
type iNetwork struct{ Vtbl *iNetworkVtbl }
|
||||||
|
|
||||||
|
func (n *iNetwork) Release() {
|
||||||
|
syscall.SyscallN(n.Vtbl.Release, uintptr(unsafe.Pointer(n)))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (n *iNetwork) GetCategory() (networkCategory, error) {
|
||||||
|
var c networkCategory
|
||||||
|
r1, _, _ := syscall.SyscallN(n.Vtbl.GetCategory,
|
||||||
|
uintptr(unsafe.Pointer(n)), uintptr(unsafe.Pointer(&c)),
|
||||||
|
)
|
||||||
|
if hr := hresult(r1); hr.failed() {
|
||||||
|
return 0, fmt.Errorf("INetwork.GetCategory: %s", hr)
|
||||||
|
}
|
||||||
|
return c, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (n *iNetwork) SetCategory(c networkCategory) error {
|
||||||
|
r1, _, _ := syscall.SyscallN(n.Vtbl.SetCategory,
|
||||||
|
uintptr(unsafe.Pointer(n)), uintptr(int32(c)),
|
||||||
|
)
|
||||||
|
if hr := hresult(r1); hr.failed() {
|
||||||
|
return fmt.Errorf("INetwork.SetCategory: %s", hr)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// coInit initializes COM for the current OS thread. The returned function must
|
||||||
|
// be deferred to balance a successful init. RPC_E_CHANGED_MODE means COM is
|
||||||
|
// already initialized in a different mode on this thread, which is still fine
|
||||||
|
// for our calls but we must not Uninitialize in that case.
|
||||||
|
func coInit() (func(), error) {
|
||||||
|
err := windows.CoInitializeEx(0, windows.COINIT_MULTITHREADED)
|
||||||
|
if err == nil {
|
||||||
|
return windows.CoUninitialize, nil
|
||||||
|
}
|
||||||
|
if e, ok := err.(syscall.Errno); ok {
|
||||||
|
switch uint32(e) {
|
||||||
|
case hrSFALSE:
|
||||||
|
return windows.CoUninitialize, nil
|
||||||
|
case hrRPCEChangedMode:
|
||||||
|
return func() {}, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("CoInitializeEx: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func createNetworkListManager() (*iNetworkListManager, error) {
|
||||||
|
var nlm *iNetworkListManager
|
||||||
|
r1, _, _ := procCoCreateInstance.Call(
|
||||||
|
uintptr(unsafe.Pointer(&clsidNetworkListManager)),
|
||||||
|
0,
|
||||||
|
uintptr(clsCtxAll),
|
||||||
|
uintptr(unsafe.Pointer(&iidINetworkListManager)),
|
||||||
|
uintptr(unsafe.Pointer(&nlm)),
|
||||||
|
)
|
||||||
|
if hr := hresult(r1); hr.failed() {
|
||||||
|
return nil, fmt.Errorf("CoCreateInstance(NetworkListManager): %s", hr)
|
||||||
|
}
|
||||||
|
return nlm, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// setNetworkCategory locates the network connection bound to adapterGUID and
|
||||||
|
// sets the category of its parent network. Returns errAdapterNotFound if the
|
||||||
|
// adapter is not yet visible in the NLM enumeration.
|
||||||
|
func setNetworkCategory(adapterGUID windows.GUID, cat networkCategory) error {
|
||||||
|
deinit, err := coInit()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer deinit()
|
||||||
|
|
||||||
|
nlm, err := createNetworkListManager()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer nlm.Release()
|
||||||
|
|
||||||
|
enum, err := nlm.GetNetworkConnections()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer enum.Release()
|
||||||
|
|
||||||
|
for {
|
||||||
|
conn, err := enum.Next()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if conn == nil {
|
||||||
|
return errAdapterNotFound
|
||||||
|
}
|
||||||
|
|
||||||
|
guid, err := conn.GetAdapterId()
|
||||||
|
if err != nil || guid != adapterGUID {
|
||||||
|
conn.Release()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
net, err := conn.GetNetwork()
|
||||||
|
conn.Release()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
err = net.SetCategory(cat)
|
||||||
|
net.Release()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// applyNetworkCategory polls until the wintun adapter shows up in the NLM
|
||||||
|
// enumeration, then sets the category. Intended to run in its own goroutine.
|
||||||
|
func applyNetworkCategory(l *slog.Logger, adapterGUID windows.GUID, cat networkCategory) {
|
||||||
|
// COM Init/Uninit must be paired on the same OS thread.
|
||||||
|
runtime.LockOSThread()
|
||||||
|
defer runtime.UnlockOSThread()
|
||||||
|
|
||||||
|
const (
|
||||||
|
attempts = 30
|
||||||
|
interval = 500 * time.Millisecond
|
||||||
|
)
|
||||||
|
for i := 0; i < attempts; i++ {
|
||||||
|
err := setNetworkCategory(adapterGUID, cat)
|
||||||
|
if err == nil {
|
||||||
|
l.Info("Set Windows network category", "category", cat.String())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !errors.Is(err, errAdapterNotFound) {
|
||||||
|
l.Warn("Failed to set Windows network category", "error", err, "category", cat.String())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(interval)
|
||||||
|
}
|
||||||
|
l.Warn("Gave up waiting for adapter to appear in NLM enumeration; network category not set",
|
||||||
|
"category", cat.String(),
|
||||||
|
"waited", time.Duration(attempts)*interval,
|
||||||
|
)
|
||||||
|
}
|
||||||
@@ -0,0 +1,109 @@
|
|||||||
|
//go:build !e2e_testing
|
||||||
|
// +build !e2e_testing
|
||||||
|
|
||||||
|
package overlay
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func Test_parseNetworkCategory(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
in string
|
||||||
|
wantCat networkCategory
|
||||||
|
wantApply bool
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"", 0, false, false},
|
||||||
|
{"unset", 0, false, false},
|
||||||
|
{" UNSET ", 0, false, false},
|
||||||
|
{"private", networkCategoryPrivate, true, false},
|
||||||
|
{"Private", networkCategoryPrivate, true, false},
|
||||||
|
{" PRIVATE ", networkCategoryPrivate, true, false},
|
||||||
|
{"public", networkCategoryPublic, true, false},
|
||||||
|
{"PUBLIC", networkCategoryPublic, true, false},
|
||||||
|
{"domain", networkCategoryDomainAuthenticated, true, false},
|
||||||
|
{"DomainAuthenticated", networkCategoryDomainAuthenticated, true, false},
|
||||||
|
{"garbage", 0, false, true},
|
||||||
|
{"privates", 0, false, true},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
cat, apply, err := parseNetworkCategory(tc.in)
|
||||||
|
if (err != nil) != tc.wantErr {
|
||||||
|
t.Errorf("parseNetworkCategory(%q) err=%v, wantErr=%v", tc.in, err, tc.wantErr)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if cat != tc.wantCat || apply != tc.wantApply {
|
||||||
|
t.Errorf("parseNetworkCategory(%q) = (%v, %v), want (%v, %v)", tc.in, cat, apply, tc.wantCat, tc.wantApply)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test_NLM_round_trip exercises every COM call path used by setNetworkCategory
|
||||||
|
// without mutating the host's network state. It validates the CLSID/IID
|
||||||
|
// constants and every vtable index by enumerating connections, fetching the
|
||||||
|
// adapter id and parent network, reading the current category, and writing it
|
||||||
|
// back unchanged.
|
||||||
|
//
|
||||||
|
// Requires Windows but does not require admin or the wintun driver. Skips if
|
||||||
|
// no network connections are available (unlikely outside of an isolated
|
||||||
|
// container).
|
||||||
|
func Test_NLM_round_trip(t *testing.T) {
|
||||||
|
deinit, err := coInit()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("coInit: %v", err)
|
||||||
|
}
|
||||||
|
defer deinit()
|
||||||
|
|
||||||
|
nlm, err := createNetworkListManager()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("createNetworkListManager: %v", err)
|
||||||
|
}
|
||||||
|
defer nlm.Release()
|
||||||
|
|
||||||
|
enum, err := nlm.GetNetworkConnections()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetNetworkConnections: %v", err)
|
||||||
|
}
|
||||||
|
defer enum.Release()
|
||||||
|
|
||||||
|
saw := 0
|
||||||
|
for {
|
||||||
|
conn, err := enum.Next()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("EnumNetworkConnections.Next: %v", err)
|
||||||
|
}
|
||||||
|
if conn == nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
saw++
|
||||||
|
|
||||||
|
if _, err := conn.GetAdapterId(); err != nil {
|
||||||
|
conn.Release()
|
||||||
|
t.Fatalf("INetworkConnection.GetAdapterId: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
net, err := conn.GetNetwork()
|
||||||
|
conn.Release()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("INetworkConnection.GetNetwork: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cat, err := net.GetCategory()
|
||||||
|
if err != nil {
|
||||||
|
net.Release()
|
||||||
|
t.Fatalf("INetwork.GetCategory: %v", err)
|
||||||
|
}
|
||||||
|
// Set to the current value so the host's NLM state is unchanged but
|
||||||
|
// SetCategory's vtable slot is still validated end-to-end.
|
||||||
|
if err := net.SetCategory(cat); err != nil {
|
||||||
|
net.Release()
|
||||||
|
t.Fatalf("INetwork.SetCategory(%v): %v", cat, err)
|
||||||
|
}
|
||||||
|
net.Release()
|
||||||
|
}
|
||||||
|
|
||||||
|
if saw == 0 {
|
||||||
|
t.Skip("no NLM network connections available; skipping round-trip")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -40,6 +40,7 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
|
|||||||
|
|
||||||
err := t.reload(c, true)
|
err := t.reload(c, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
_ = file.Close()
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
//go:build (amd64 || arm64) && !e2e_testing
|
||||||
|
// +build amd64 arm64
|
||||||
|
// +build !e2e_testing
|
||||||
|
|
||||||
|
package overlay
|
||||||
|
|
||||||
|
import (
|
||||||
|
"log/slog"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/wfp"
|
||||||
|
)
|
||||||
|
|
||||||
|
// installInterfaceBypass installs a WFP PERMIT filter scoped to the wintun interface LUID so inbound traffic on the
|
||||||
|
// nebula adapter bypasses Windows Defender Firewall.
|
||||||
|
func installInterfaceBypass(l *slog.Logger, luid uint64) closer {
|
||||||
|
s, err := wfp.PermitInterface(luid)
|
||||||
|
if err != nil {
|
||||||
|
l.Warn("Failed to install WFP bypass filters on nebula interface", "error", err)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
l.Info("Installed WFP filters bypassing Windows Defender Firewall on nebula interface")
|
||||||
|
return s
|
||||||
|
}
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
//go:build !e2e_testing
|
||||||
|
// +build !e2e_testing
|
||||||
|
|
||||||
|
package overlay
|
||||||
|
|
||||||
|
import "log/slog"
|
||||||
|
|
||||||
|
// installInterfaceBypass is a no-op on windows-386 because we don't currently build for it.
|
||||||
|
func installInterfaceBypass(_ *slog.Logger, _ uint64) closer {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+86
-28
@@ -23,7 +23,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
io.ReadWriteCloser
|
f *os.File
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
DefaultMTU int
|
DefaultMTU int
|
||||||
@@ -31,9 +31,6 @@ type tun struct {
|
|||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||||
linkAddr *netroute.LinkAddr
|
linkAddr *netroute.LinkAddr
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
|
|
||||||
// cache out buffer since we need to prepend 4 bytes for tun metadata
|
|
||||||
out []byte
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type ifReq struct {
|
type ifReq struct {
|
||||||
@@ -124,7 +121,7 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, _ bool) (*t
|
|||||||
}
|
}
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
ReadWriteCloser: os.NewFile(uintptr(fd), ""),
|
f: os.NewFile(uintptr(fd), ""),
|
||||||
Device: name,
|
Device: name,
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
DefaultMTU: c.GetInt("tun.mtu", DefaultMTU),
|
DefaultMTU: c.GetInt("tun.mtu", DefaultMTU),
|
||||||
@@ -158,8 +155,8 @@ func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, e
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Close() error {
|
func (t *tun) Close() error {
|
||||||
if t.ReadWriteCloser != nil {
|
if t.f != nil {
|
||||||
return t.ReadWriteCloser.Close()
|
return t.f.Close()
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -502,42 +499,103 @@ func delRoute(prefix netip.Prefix, gateway netroute.Addr) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// tunWritev and tunReadv are linkname'd to x/sys/unix's libc-routed writev/readv stubs so the
|
||||||
|
// calls go through libSystem's pinned trampoline. A raw syscall.Syscall(SYS_WRITEV/SYS_READV, ...)
|
||||||
|
// on darwin/arm64 emits an SVC #0x80 trap (see $GOROOT/src/syscall/asm_darwin_arm64.s), the path
|
||||||
|
// Apple keeps warning they will eventually disallow. We pull the low-level stubs instead of calling
|
||||||
|
// unix.Writev/unix.Readv because those take [][]byte and rebuild the []Iovec every call, which
|
||||||
|
// heap-allocates the header; linkname'ing the stubs lets us hand them our own stack-allocated
|
||||||
|
// iovecs. See golang/go#78049.
|
||||||
|
|
||||||
|
//go:linkname tunWritev golang.org/x/sys/unix.writev
|
||||||
|
//go:noescape
|
||||||
|
func tunWritev(fd int, iovecs []unix.Iovec) (n int, err error)
|
||||||
|
|
||||||
|
//go:linkname tunReadv golang.org/x/sys/unix.readv
|
||||||
|
//go:noescape
|
||||||
|
func tunReadv(fd int, iovecs []unix.Iovec) (n int, err error)
|
||||||
|
|
||||||
|
// Read pulls one IP packet off the utun device, scattering the 4 byte protocol header away from
|
||||||
|
// the packet so the payload lands directly in to.
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
func (t *tun) Read(to []byte) (int, error) {
|
||||||
buf := make([]byte, len(to)+4)
|
var head [4]byte
|
||||||
|
|
||||||
n, err := t.ReadWriteCloser.Read(buf)
|
rc, err := t.f.SyscallConn()
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
copy(to, buf[4:])
|
var n int
|
||||||
return n - 4, err
|
var callErr error
|
||||||
|
err = rc.Read(func(fd uintptr) bool {
|
||||||
|
iovecs := []unix.Iovec{
|
||||||
|
{Base: &head[0], Len: 4},
|
||||||
|
{Base: &to[0], Len: uint64(len(to))},
|
||||||
|
}
|
||||||
|
n, callErr = tunReadv(int(fd), iovecs)
|
||||||
|
if errno, ok := callErr.(syscall.Errno); ok && errno.Temporary() {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
if callErr != nil {
|
||||||
|
return 0, callErr
|
||||||
|
}
|
||||||
|
if n < 4 {
|
||||||
|
return 0, nil
|
||||||
|
}
|
||||||
|
return n - 4, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Write is only valid for single threaded use
|
// Write pushes one IP packet onto the utun device.
|
||||||
func (t *tun) Write(from []byte) (int, error) {
|
func (t *tun) Write(from []byte) (int, error) {
|
||||||
buf := t.out
|
|
||||||
if cap(buf) < len(from)+4 {
|
|
||||||
buf = make([]byte, len(from)+4)
|
|
||||||
t.out = buf
|
|
||||||
}
|
|
||||||
buf = buf[:len(from)+4]
|
|
||||||
|
|
||||||
if len(from) == 0 {
|
if len(from) == 0 {
|
||||||
return 0, syscall.EIO
|
return 0, syscall.EIO
|
||||||
}
|
}
|
||||||
|
|
||||||
// Determine the IP Family for the NULL L2 Header
|
|
||||||
ipVer := from[0] >> 4
|
ipVer := from[0] >> 4
|
||||||
if ipVer == 4 {
|
var head [4]byte
|
||||||
buf[3] = syscall.AF_INET
|
switch ipVer {
|
||||||
} else if ipVer == 6 {
|
case 4:
|
||||||
buf[3] = syscall.AF_INET6
|
head[3] = syscall.AF_INET
|
||||||
} else {
|
case 6:
|
||||||
|
head[3] = syscall.AF_INET6
|
||||||
|
default:
|
||||||
return 0, fmt.Errorf("unable to determine IP version from packet")
|
return 0, fmt.Errorf("unable to determine IP version from packet")
|
||||||
}
|
}
|
||||||
|
|
||||||
copy(buf[4:], from)
|
// Grab rc as a local so the compiler can devirtualize the call and keep the closure on the stack.
|
||||||
|
rc, err := t.f.SyscallConn()
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
n, err := t.ReadWriteCloser.Write(buf)
|
var n int
|
||||||
return n - 4, err
|
var callErr error
|
||||||
|
err = rc.Write(func(fd uintptr) bool {
|
||||||
|
iovecs := []unix.Iovec{
|
||||||
|
{Base: &head[0], Len: 4},
|
||||||
|
{Base: &from[0], Len: uint64(len(from))},
|
||||||
|
}
|
||||||
|
n, callErr = tunWritev(int(fd), iovecs)
|
||||||
|
// Type-assert to syscall.Errno so the EAGAIN/EWOULDBLOCK/EINTR check doesn't box the errno
|
||||||
|
// constants into error interfaces on every call.
|
||||||
|
if errno, ok := callErr.(syscall.Errno); ok && errno.Temporary() {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
if callErr != nil {
|
||||||
|
return 0, callErr
|
||||||
|
}
|
||||||
|
|
||||||
|
return n - 4, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Networks() []netip.Prefix {
|
func (t *tun) Networks() []netip.Prefix {
|
||||||
|
|||||||
@@ -659,7 +659,6 @@ func addRoute(prefix netip.Prefix, gateway netroute.Addr) error {
|
|||||||
return fmt.Errorf("failed to create route.RouteMessage for change: %w", err)
|
return fmt.Errorf("failed to create route.RouteMessage for change: %w", err)
|
||||||
}
|
}
|
||||||
_, err = unix.Write(sock, data[:])
|
_, err = unix.Write(sock, data[:])
|
||||||
fmt.Println("DOING CHANGE")
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return fmt.Errorf("failed to write route.RouteMessage to socket: %w", err)
|
return fmt.Errorf("failed to write route.RouteMessage to socket: %w", err)
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
@@ -33,6 +34,12 @@ func newTun(_ *config.C, _ *slog.Logger, _ []netip.Prefix, _ bool) (*tun, error)
|
|||||||
}
|
}
|
||||||
|
|
||||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
|
if err := unix.SetNonblock(deviceFd, true); err != nil {
|
||||||
|
// We own the fd from the moment it is handed to us, same as the reload error path below
|
||||||
|
_ = unix.Close(deviceFd)
|
||||||
|
return nil, fmt.Errorf("failed to set the tun fd to non-blocking mode: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
file := os.NewFile(uintptr(deviceFd), "/dev/tun")
|
file := os.NewFile(uintptr(deviceFd), "/dev/tun")
|
||||||
t := &tun{
|
t := &tun{
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
@@ -42,6 +49,7 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
|
|||||||
|
|
||||||
err := t.reload(c, true)
|
err := t.reload(c, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
_ = file.Close()
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ package overlay
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
@@ -483,8 +484,17 @@ func (t *tun) addIPs(link netlink.Link) error {
|
|||||||
//iterate over remainder, remove whoever shouldn't be there
|
//iterate over remainder, remove whoever shouldn't be there
|
||||||
al, err := netlink.AddrList(link, netlink.FAMILY_ALL)
|
al, err := netlink.AddrList(link, netlink.FAMILY_ALL)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
//RTM_GETADDR dumps the whole system, so any concurrent address change
|
||||||
|
//interrupts it - including the kernel's async tentative->preferred
|
||||||
|
//flip of an IPv6 address the AddrReplace calls above just added,
|
||||||
|
//which makes this a race against our own setup. Partial results are
|
||||||
|
//still returned; the worst case is a stale address surviving until
|
||||||
|
//the next config reload, which beats failing startup over it.
|
||||||
|
if !errors.Is(err, netlink.ErrDumpInterrupted) {
|
||||||
return fmt.Errorf("failed to get tun address list: %s", err)
|
return fmt.Errorf("failed to get tun address list: %s", err)
|
||||||
}
|
}
|
||||||
|
t.l.Warn("tun address list dump was interrupted, stale addresses may remain")
|
||||||
|
}
|
||||||
|
|
||||||
for i := range al {
|
for i := range al {
|
||||||
if hasNetlinkAddr(newAddrs, al[i]) {
|
if hasNetlinkAddr(newAddrs, al[i]) {
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user