mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-16 01:37:00 +02:00
Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| e9357ff426 |
@@ -1,116 +0,0 @@
|
|||||||
name: Code-sign Windows binaries
|
|
||||||
description: >
|
|
||||||
Sign every .exe under a given path in place via the DefinedNet code-signer
|
|
||||||
Lambda. If `role` or `bucket` is empty, logs a notice and skips signing so
|
|
||||||
forks and dev branches without AWS access still produce usable builds.
|
|
||||||
|
|
||||||
inputs:
|
|
||||||
path:
|
|
||||||
description: "Directory whose .exe files should be signed in place"
|
|
||||||
required: true
|
|
||||||
role:
|
|
||||||
description: "IAM role ARN to assume via OIDC; empty disables signing"
|
|
||||||
required: false
|
|
||||||
default: ""
|
|
||||||
bucket:
|
|
||||||
description: "S3 staging bucket the code-signer Lambda reads from; empty disables signing"
|
|
||||||
required: false
|
|
||||||
default: ""
|
|
||||||
region:
|
|
||||||
description: "AWS region for the role and Lambda"
|
|
||||||
required: false
|
|
||||||
default: "us-east-2"
|
|
||||||
function-name:
|
|
||||||
description: "Code-signer Lambda function name"
|
|
||||||
required: false
|
|
||||||
default: "code-signer"
|
|
||||||
key-prefix:
|
|
||||||
description: "S3 key prefix 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
|
|
||||||
@@ -0,0 +1,34 @@
|
|||||||
|
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@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.26'
|
go-version: '1.25'
|
||||||
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@v7
|
uses: actions/upload-artifact@v6
|
||||||
with:
|
with:
|
||||||
name: linux-latest
|
name: linux-latest
|
||||||
path: release
|
path: release
|
||||||
@@ -32,15 +32,12 @@ 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@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.26'
|
go-version: '1.25'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
@@ -57,15 +54,8 @@ jobs:
|
|||||||
mkdir build\dist\windows
|
mkdir build\dist\windows
|
||||||
mv dist\windows\wintun build\dist\windows\
|
mv dist\windows\wintun build\dist\windows\
|
||||||
|
|
||||||
- name: Code-sign
|
|
||||||
uses: ./.github/actions/code-sign
|
|
||||||
with:
|
|
||||||
path: build
|
|
||||||
role: ${{ secrets.DEFINED_CODE_SIGNER_ROLE }}
|
|
||||||
bucket: ${{ secrets.DEFINED_CODE_SIGNER_BUCKET }}
|
|
||||||
|
|
||||||
- name: Upload artifacts
|
- name: Upload artifacts
|
||||||
uses: actions/upload-artifact@v7
|
uses: actions/upload-artifact@v6
|
||||||
with:
|
with:
|
||||||
name: windows-latest
|
name: windows-latest
|
||||||
path: build
|
path: build
|
||||||
@@ -76,16 +66,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@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.26'
|
go-version: '1.25'
|
||||||
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@v7
|
uses: Apple-Actions/import-codesign-certs@v6
|
||||||
with:
|
with:
|
||||||
p12-file-base64: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_P12_BASE64 }}
|
p12-file-base64: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_P12_BASE64 }}
|
||||||
p12-password: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_PASSWORD }}
|
p12-password: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_PASSWORD }}
|
||||||
@@ -114,7 +104,7 @@ jobs:
|
|||||||
fi
|
fi
|
||||||
|
|
||||||
- name: Upload artifacts
|
- name: Upload artifacts
|
||||||
uses: actions/upload-artifact@v7
|
uses: actions/upload-artifact@v6
|
||||||
with:
|
with:
|
||||||
name: darwin-latest
|
name: darwin-latest
|
||||||
path: ./release/*
|
path: ./release/*
|
||||||
@@ -134,25 +124,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@v7
|
uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
uses: actions/download-artifact@v8
|
uses: actions/download-artifact@v7
|
||||||
with:
|
with:
|
||||||
name: linux-latest
|
name: linux-latest
|
||||||
path: artifacts
|
path: artifacts
|
||||||
|
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
uses: docker/login-action@v4
|
uses: docker/login-action@v3
|
||||||
with:
|
with:
|
||||||
username: ${{ vars.DOCKERHUB_USERNAME }}
|
username: ${{ vars.DOCKERHUB_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_TOKEN }}
|
password: ${{ secrets.DOCKERHUB_TOKEN }}
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
uses: docker/setup-buildx-action@v4
|
uses: docker/setup-buildx-action@v3
|
||||||
|
|
||||||
- name: Build and push images
|
- name: Build and push images
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
@@ -163,20 +153,17 @@ 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 \
|
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}"
|
||||||
--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@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
uses: actions/download-artifact@v8
|
uses: actions/download-artifact@v7
|
||||||
with:
|
with:
|
||||||
path: artifacts
|
path: artifacts
|
||||||
|
|
||||||
|
|||||||
@@ -14,27 +14,19 @@ on:
|
|||||||
- 'go.sum'
|
- 'go.sum'
|
||||||
jobs:
|
jobs:
|
||||||
|
|
||||||
smoke-extra-libvirt:
|
smoke-extra:
|
||||||
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
||||||
name: ${{ matrix.target }}
|
name: Run extra smoke tests
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
strategy:
|
|
||||||
fail-fast: false
|
|
||||||
matrix:
|
|
||||||
target:
|
|
||||||
- freebsd-amd64
|
|
||||||
- openbsd-amd64
|
|
||||||
- netbsd-amd64
|
|
||||||
- linux-amd64-ipv6disable
|
|
||||||
env:
|
env:
|
||||||
VAGRANT_DEFAULT_PROVIDER: libvirt
|
VAGRANT_DEFAULT_PROVIDER: libvirt
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.26'
|
go-version: '1.25'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: add hashicorp source
|
- name: add hashicorp source
|
||||||
@@ -48,85 +40,28 @@ jobs:
|
|||||||
sudo chmod 666 /var/run/libvirt/libvirt-sock
|
sudo chmod 666 /var/run/libvirt/libvirt-sock
|
||||||
vagrant plugin install vagrant-libvirt
|
vagrant plugin install vagrant-libvirt
|
||||||
|
|
||||||
- name: ${{ matrix.target }}
|
- name: freebsd-amd64
|
||||||
run: make smoke-vagrant/${{ matrix.target }}
|
run: make smoke-vagrant/freebsd-amd64
|
||||||
|
|
||||||
timeout-minutes: 30
|
- name: openbsd-amd64
|
||||||
|
run: make smoke-vagrant/openbsd-amd64
|
||||||
|
|
||||||
# linux-386 needs VirtualBox, which conflicts with KVM/libvirt -- isolated job.
|
- name: netbsd-amd64
|
||||||
smoke-extra-virtualbox:
|
run: make smoke-vagrant/netbsd-amd64
|
||||||
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
|
||||||
name: linux-386
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
env:
|
|
||||||
VAGRANT_DEFAULT_PROVIDER: virtualbox
|
|
||||||
steps:
|
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- name: linux-amd64-ipv6disable
|
||||||
|
run: make smoke-vagrant/linux-amd64-ipv6disable
|
||||||
|
|
||||||
- uses: actions/setup-go@v7
|
# linux-386 runs last because it requires disabling KVM to use VirtualBox,
|
||||||
with:
|
# which prevents libvirt (used by the other tests) from working after this point.
|
||||||
go-version: '1.26'
|
- name: install virtualbox for i386 test
|
||||||
check-latest: true
|
|
||||||
|
|
||||||
- name: add hashicorp source
|
|
||||||
run: wget -O- https://apt.releases.hashicorp.com/gpg | gpg --dearmor | sudo tee /usr/share/keyrings/hashicorp-archive-keyring.gpg && echo "deb [signed-by=/usr/share/keyrings/hashicorp-archive-keyring.gpg] https://apt.releases.hashicorp.com $(lsb_release -cs) main" | sudo tee /etc/apt/sources.list.d/hashicorp.list
|
|
||||||
|
|
||||||
- name: install vagrant and virtualbox
|
|
||||||
run: |
|
run: |
|
||||||
sudo apt-get update && sudo apt-get install -y vagrant virtualbox
|
sudo apt-get install -y virtualbox
|
||||||
sudo rmmod kvm_amd kvm_intel kvm 2>/dev/null || true
|
sudo rmmod kvm_amd kvm_intel kvm 2>/dev/null || true
|
||||||
|
|
||||||
- name: linux-386
|
- name: linux-386
|
||||||
|
env:
|
||||||
|
VAGRANT_DEFAULT_PROVIDER: virtualbox
|
||||||
run: make smoke-vagrant/linux-386
|
run: make smoke-vagrant/linux-386
|
||||||
|
|
||||||
timeout-minutes: 30
|
timeout-minutes: 30
|
||||||
|
|
||||||
smoke-windows:
|
|
||||||
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
|
||||||
name: Run windows smoke test
|
|
||||||
runs-on: windows-latest
|
|
||||||
steps:
|
|
||||||
|
|
||||||
- uses: actions/checkout@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
|
|
||||||
|
|||||||
+25
-13
@@ -18,26 +18,38 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.26'
|
go-version: '1.25'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Smoke Docker
|
- name: build
|
||||||
run: make smoke-docker
|
run: make bin-docker CGO_ENABLED=1 BUILD_ARGS=-race
|
||||||
|
|
||||||
- name: Smoke Docker IPv6 overlay
|
- name: setup docker image
|
||||||
run: make smoke-docker-ipv6
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: ./build.sh
|
||||||
|
|
||||||
- name: Smoke Relay Docker
|
- name: run smoke
|
||||||
run: make smoke-relay-docker
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: ./smoke.sh
|
||||||
|
|
||||||
- name: Smoke Docker boringcrypto
|
- name: setup relay docker image
|
||||||
run: make boringcrypto smoke-docker
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: ./build-relay.sh
|
||||||
|
|
||||||
- name: Smoke Docker fips140
|
- name: run smoke relay
|
||||||
run: make fips140-all GOALS=smoke-docker
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: ./smoke-relay.sh
|
||||||
|
|
||||||
|
- name: setup docker image for P256
|
||||||
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: NAME="smoke-p256" CURVE=P256 ./build.sh
|
||||||
|
|
||||||
|
- name: run smoke-p256
|
||||||
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: NAME="smoke-p256" ./smoke.sh
|
||||||
|
|
||||||
timeout-minutes: 10
|
timeout-minutes: 10
|
||||||
|
|||||||
@@ -5,19 +5,6 @@ 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
|
||||||
@@ -44,24 +31,24 @@ LIGHTHOUSE_IP="203.0.113.2"
|
|||||||
../genconfig.sh >lighthouse1.yml
|
../genconfig.sh >lighthouse1.yml
|
||||||
|
|
||||||
HOST="host2" \
|
HOST="host2" \
|
||||||
LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
||||||
../genconfig.sh >host2.yml
|
../genconfig.sh >host2.yml
|
||||||
|
|
||||||
HOST="host3" \
|
HOST="host3" \
|
||||||
LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="192.168.100.1 $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="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="192.168.100.1 $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 "$LIGHTHOUSE_NIP/24"
|
../../../../nebula-cert sign -name "lighthouse1" -groups "lighthouse,lighthouse1" -ip "192.168.100.1/24"
|
||||||
../../../../nebula-cert sign -name "host2" -groups "host,host2" -ip "$HOST2_NIP/24"
|
../../../../nebula-cert sign -name "host2" -groups "host,host2" -ip "192.168.100.2/24"
|
||||||
../../../../nebula-cert sign -name "host3" -groups "host,host3" -ip "$HOST3_NIP/24"
|
../../../../nebula-cert sign -name "host3" -groups "host,host3" -ip "192.168.100.3/24"
|
||||||
../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "$HOST4_NIP/24"
|
../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "192.168.100.4/24"
|
||||||
)
|
)
|
||||||
|
|
||||||
docker build -t "nebula:${NAME:-smoke}" .
|
docker build -t "nebula:${NAME:-smoke}" .
|
||||||
|
|||||||
@@ -1,272 +0,0 @@
|
|||||||
#!/usr/bin/env pwsh
|
|
||||||
# Windows smoke test for the nebula tun + UDP + NLM code paths.
|
|
||||||
#
|
|
||||||
# Topology:
|
|
||||||
# - lighthouse runs natively on the Windows host (wintun + windows UDP)
|
|
||||||
# - peer runs inside WSL2 (Linux build of nebula, /dev/net/tun)
|
|
||||||
#
|
|
||||||
# WSL2 gives us a real netns boundary so the loopback fast-path on Windows
|
|
||||||
# does not short-circuit the overlay -- when WSL pings the lighthouse VPN IP,
|
|
||||||
# Linux has no idea that IP is local to the Windows host, so the packet is
|
|
||||||
# forced through nebula. Same in reverse.
|
|
||||||
|
|
||||||
$ErrorActionPreference = 'Stop'
|
|
||||||
|
|
||||||
# wsl.exe emits UTF-16 LE by default which PowerShell reads as bytes, mangling
|
|
||||||
# every captured string. WSL_UTF8 makes wsl.exe emit UTF-8 instead.
|
|
||||||
$env:WSL_UTF8 = '1'
|
|
||||||
|
|
||||||
$RepoRoot = Resolve-Path "$PSScriptRoot\..\..\.."
|
|
||||||
$Nebula = Join-Path $RepoRoot 'nebula.exe'
|
|
||||||
$NebulaCert = Join-Path $RepoRoot 'nebula-cert.exe'
|
|
||||||
$NebulaLinux = Join-Path $RepoRoot 'build\linux-amd64\nebula'
|
|
||||||
|
|
||||||
if (-not (Test-Path $Nebula)) { throw "missing $Nebula; run 'make bin-windows' first" }
|
|
||||||
if (-not (Test-Path $NebulaCert)) { throw "missing $NebulaCert; run 'make bin-windows' first" }
|
|
||||||
if (-not (Test-Path $NebulaLinux)) { throw "missing $NebulaLinux; build the linux nebula first" }
|
|
||||||
|
|
||||||
# Matches the distro installed by Vampire/setup-wsl in smoke-extra.yml.
|
|
||||||
$Distro = 'Ubuntu-24.04'
|
|
||||||
$listed = (wsl --list --quiet 2>$null) -join "`n"
|
|
||||||
if ($listed -notmatch [regex]::Escape($Distro)) {
|
|
||||||
throw "WSL distro $Distro not registered. Got: $listed"
|
|
||||||
}
|
|
||||||
Write-Host "Using WSL distro: $Distro"
|
|
||||||
|
|
||||||
# Windows host as seen from inside WSL: WSL's default-route gateway. We extract
|
|
||||||
# it with a regex rather than awk fields so PowerShell does not eat any '$N'
|
|
||||||
# tokens, and tabs/double-spaces in `ip route` output do not confuse a cut.
|
|
||||||
$ipCmd = 'ip route show default | grep -oE "([0-9]+\.){3}[0-9]+" | head -1'
|
|
||||||
$WindowsIp = (wsl -d $Distro -- bash -c $ipCmd).Trim()
|
|
||||||
if (-not $WindowsIp) { throw "could not determine Windows host IP from WSL" }
|
|
||||||
Write-Host "Windows host IP from WSL: $WindowsIp"
|
|
||||||
|
|
||||||
$WorkDir = Join-Path $env:TEMP 'nebula-smoke-windows'
|
|
||||||
if (Test-Path $WorkDir) { Remove-Item -Recurse -Force $WorkDir }
|
|
||||||
New-Item -ItemType Directory -Path $WorkDir | Out-Null
|
|
||||||
|
|
||||||
$WslDir = '/tmp/nebula-smoke'
|
|
||||||
wsl -d $Distro -- bash -c "rm -rf $WslDir && mkdir -p $WslDir" | Out-Null
|
|
||||||
|
|
||||||
$DevName = 'nebula-smoke'
|
|
||||||
$Ip1 = '192.168.241.1'
|
|
||||||
$Ip2 = '192.168.241.2'
|
|
||||||
$Port = 4242
|
|
||||||
|
|
||||||
& $NebulaCert ca -name 'smoke-ca' -out-crt "$WorkDir\ca.crt" -out-key "$WorkDir\ca.key"
|
|
||||||
if ($LASTEXITCODE -ne 0) { throw "nebula-cert ca failed (exit $LASTEXITCODE)" }
|
|
||||||
|
|
||||||
& $NebulaCert sign -name 'lighthouse' -networks "$Ip1/24" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\lighthouse.crt" -out-key "$WorkDir\lighthouse.key"
|
|
||||||
if ($LASTEXITCODE -ne 0) { throw "nebula-cert sign lighthouse failed (exit $LASTEXITCODE)" }
|
|
||||||
|
|
||||||
& $NebulaCert sign -name 'peer' -networks "$Ip2/24" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\peer.crt" -out-key "$WorkDir\peer.key"
|
|
||||||
if ($LASTEXITCODE -ne 0) { throw "nebula-cert sign peer failed (exit $LASTEXITCODE)" }
|
|
||||||
|
|
||||||
# Windows lighthouse config.
|
|
||||||
@"
|
|
||||||
pki:
|
|
||||||
ca: $WorkDir\ca.crt
|
|
||||||
cert: $WorkDir\lighthouse.crt
|
|
||||||
key: $WorkDir\lighthouse.key
|
|
||||||
static_host_map: {}
|
|
||||||
lighthouse:
|
|
||||||
am_lighthouse: true
|
|
||||||
interval: 60
|
|
||||||
hosts: []
|
|
||||||
listen:
|
|
||||||
host: 0.0.0.0
|
|
||||||
port: $Port
|
|
||||||
tun:
|
|
||||||
disabled: false
|
|
||||||
dev: $DevName
|
|
||||||
drop_local_broadcast: false
|
|
||||||
drop_multicast: false
|
|
||||||
tx_queue: 500
|
|
||||||
mtu: 1300
|
|
||||||
network_category: private
|
|
||||||
logging:
|
|
||||||
level: info
|
|
||||||
format: text
|
|
||||||
firewall:
|
|
||||||
outbound_action: drop
|
|
||||||
inbound_action: drop
|
|
||||||
conntrack:
|
|
||||||
tcp_timeout: 12m
|
|
||||||
udp_timeout: 3m
|
|
||||||
default_timeout: 10m
|
|
||||||
outbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
inbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
"@ | Out-File -FilePath "$WorkDir\lighthouse.yml" -Encoding utf8
|
|
||||||
|
|
||||||
# WSL peer config (paths are POSIX, deliberately).
|
|
||||||
@"
|
|
||||||
pki:
|
|
||||||
ca: $WslDir/ca.crt
|
|
||||||
cert: $WslDir/peer.crt
|
|
||||||
key: $WslDir/peer.key
|
|
||||||
static_host_map:
|
|
||||||
"${Ip1}": ["${WindowsIp}:$Port"]
|
|
||||||
lighthouse:
|
|
||||||
am_lighthouse: false
|
|
||||||
interval: 60
|
|
||||||
hosts:
|
|
||||||
- "${Ip1}"
|
|
||||||
listen:
|
|
||||||
host: 0.0.0.0
|
|
||||||
port: 0
|
|
||||||
tun:
|
|
||||||
disabled: false
|
|
||||||
dev: nebula1
|
|
||||||
drop_local_broadcast: false
|
|
||||||
drop_multicast: false
|
|
||||||
tx_queue: 500
|
|
||||||
mtu: 1300
|
|
||||||
logging:
|
|
||||||
level: info
|
|
||||||
format: text
|
|
||||||
firewall:
|
|
||||||
outbound_action: drop
|
|
||||||
inbound_action: drop
|
|
||||||
conntrack:
|
|
||||||
tcp_timeout: 12m
|
|
||||||
udp_timeout: 3m
|
|
||||||
default_timeout: 10m
|
|
||||||
outbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
inbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
"@ | Out-File -FilePath "$WorkDir\peer.yml" -Encoding utf8
|
|
||||||
|
|
||||||
# Stage WSL artifacts. Convert Windows paths to WSL paths ourselves rather than
|
|
||||||
# calling `wslpath`, because PowerShell's argument-passing to external EXEs
|
|
||||||
# strips backslashes from path arguments in ways that are hard to escape around.
|
|
||||||
function ConvertTo-WslPath {
|
|
||||||
param([string]$WindowsPath)
|
|
||||||
if ($WindowsPath -notmatch '^([A-Za-z]):\\(.*)$') {
|
|
||||||
throw "cannot convert path to WSL: $WindowsPath"
|
|
||||||
}
|
|
||||||
return "/mnt/$($matches[1].ToLower())/$($matches[2].Replace('\','/'))"
|
|
||||||
}
|
|
||||||
|
|
||||||
$WslWorkDir = ConvertTo-WslPath $WorkDir
|
|
||||||
$WslNebulaPath = ConvertTo-WslPath $NebulaLinux
|
|
||||||
wsl -d $Distro -- bash -c "cp '$WslWorkDir/ca.crt' '$WslWorkDir/peer.crt' '$WslWorkDir/peer.key' '$WslWorkDir/peer.yml' $WslDir/ && cp '$WslNebulaPath' $WslDir/nebula && chmod +x $WslDir/nebula"
|
|
||||||
|
|
||||||
# Make sure WSL has tun support and /dev/net/tun is usable before starting
|
|
||||||
# nebula. Diagnostics first so a fail here points at the real problem (e.g.
|
|
||||||
# WSL1 distros do not have a real kernel and will not have tun).
|
|
||||||
Write-Host '=== WSL diagnostic ==='
|
|
||||||
wsl --version 2>&1 | Out-Host
|
|
||||||
wsl --list --verbose 2>&1 | Out-Host
|
|
||||||
wsl -d $Distro -u root -- uname -a | Out-Host
|
|
||||||
wsl -d $Distro -u root -- bash -c "modprobe tun 2>&1 || true; mkdir -p /dev/net; [ -c /dev/net/tun ] || mknod /dev/net/tun c 10 200; chmod 600 /dev/net/tun; ls -l /dev/net/tun"
|
|
||||||
if ($LASTEXITCODE -ne 0) { throw "failed to prepare /dev/net/tun in WSL (TUN support missing?)" }
|
|
||||||
|
|
||||||
# Deliberately no New-NetFirewallRule calls here -- nebula's windows_bypass_wdf
|
|
||||||
# feature is supposed to install WFP permit filters that let inbound traffic
|
|
||||||
# through Windows Defender Firewall on its own. If this smoke regresses, that
|
|
||||||
# feature regressed.
|
|
||||||
|
|
||||||
$lhOut = Join-Path $WorkDir 'lighthouse.out.log'
|
|
||||||
$lhErr = Join-Path $WorkDir 'lighthouse.err.log'
|
|
||||||
$lhProc = Start-Process -FilePath $Nebula -ArgumentList @('-config', "$WorkDir\lighthouse.yml") `
|
|
||||||
-PassThru -NoNewWindow `
|
|
||||||
-RedirectStandardOutput $lhOut `
|
|
||||||
-RedirectStandardError $lhErr
|
|
||||||
|
|
||||||
# Run nebula in WSL as root with no sudo + no shell wrapper. PowerShell's
|
|
||||||
# Start-Process arg quoting mangles `bash -c "..."` strings that contain
|
|
||||||
# spaces/redirections, so we skip bash entirely and let Start-Process do the
|
|
||||||
# stdout/stderr capture itself.
|
|
||||||
$peerOut = Join-Path $WorkDir 'peer.out.log'
|
|
||||||
$peerErr = Join-Path $WorkDir 'peer.err.log'
|
|
||||||
$peerProc = Start-Process -FilePath 'wsl' `
|
|
||||||
-ArgumentList @('-d', $Distro, '-u', 'root', '--', "$WslDir/nebula", '-config', "$WslDir/peer.yml") `
|
|
||||||
-PassThru -NoNewWindow `
|
|
||||||
-RedirectStandardOutput $peerOut `
|
|
||||||
-RedirectStandardError $peerErr
|
|
||||||
|
|
||||||
function Wait-Until {
|
|
||||||
param([scriptblock]$Predicate, [int]$TimeoutSec, [string]$What)
|
|
||||||
$deadline = (Get-Date).AddSeconds($TimeoutSec)
|
|
||||||
while ((Get-Date) -lt $deadline) {
|
|
||||||
if (& $Predicate) { return }
|
|
||||||
Start-Sleep -Milliseconds 500
|
|
||||||
}
|
|
||||||
throw "timed out waiting for: $What"
|
|
||||||
}
|
|
||||||
|
|
||||||
try {
|
|
||||||
Wait-Until -TimeoutSec 30 -What "windows wintun adapter $DevName with NetworkCategory=Private" -Predicate {
|
|
||||||
if ($lhProc.HasExited) { throw "lighthouse exited (code $($lhProc.ExitCode)) before tun was ready" }
|
|
||||||
$p = Get-NetConnectionProfile -InterfaceAlias $DevName -ErrorAction SilentlyContinue
|
|
||||||
$p -and ("$($p.NetworkCategory)" -ieq 'Private')
|
|
||||||
}
|
|
||||||
Write-Host "OK: $DevName NetworkCategory=Private"
|
|
||||||
|
|
||||||
Wait-Until -TimeoutSec 30 -What "WSL nebula1 with $Ip2" -Predicate {
|
|
||||||
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before tun was ready" }
|
|
||||||
$r = wsl -d $Distro -u root -- bash -c "ip -o addr show nebula1 2>/dev/null | grep -q 'inet $Ip2' && echo yes"
|
|
||||||
("$r").Trim() -eq 'yes'
|
|
||||||
}
|
|
||||||
Write-Host "OK: WSL nebula1 has $Ip2"
|
|
||||||
|
|
||||||
Wait-Until -TimeoutSec 30 -What "ping from WSL peer to windows lighthouse ($Ip1)" -Predicate {
|
|
||||||
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before ping succeeded" }
|
|
||||||
$r = wsl -d $Distro -u root -- bash -c "ping -c1 -W1 $Ip1 >/dev/null 2>&1 && echo OK"
|
|
||||||
("$r").Trim() -eq 'OK'
|
|
||||||
}
|
|
||||||
Write-Host "OK: WSL peer -> windows lighthouse"
|
|
||||||
|
|
||||||
Wait-Until -TimeoutSec 30 -What "ping from windows lighthouse to WSL peer ($Ip2)" -Predicate {
|
|
||||||
$null = & ping.exe -n 1 -w 1000 $Ip2
|
|
||||||
$LASTEXITCODE -eq 0
|
|
||||||
}
|
|
||||||
Write-Host "OK: windows lighthouse -> WSL peer"
|
|
||||||
|
|
||||||
Write-Host ''
|
|
||||||
Write-Host 'All smoke checks passed.'
|
|
||||||
}
|
|
||||||
catch {
|
|
||||||
Write-Host ''
|
|
||||||
Write-Host '=== lighthouse stdout ==='
|
|
||||||
Get-Content $lhOut -ErrorAction SilentlyContinue | Out-Host
|
|
||||||
Write-Host '=== lighthouse stderr ==='
|
|
||||||
Get-Content $lhErr -ErrorAction SilentlyContinue | Out-Host
|
|
||||||
Write-Host '=== peer stdout ==='
|
|
||||||
Get-Content $peerOut -ErrorAction SilentlyContinue | Out-Host
|
|
||||||
Write-Host '=== peer stderr ==='
|
|
||||||
Get-Content $peerErr -ErrorAction SilentlyContinue | Out-Host
|
|
||||||
Write-Host '=== nebula WFP filters ==='
|
|
||||||
# Dump nebula-installed filters so we can verify they got registered with
|
|
||||||
# the conditions we expect.
|
|
||||||
$wfpDump = Join-Path $WorkDir 'wfp.xml'
|
|
||||||
netsh wfp show filters file=$wfpDump 2>&1 | Out-Null
|
|
||||||
if (Test-Path $wfpDump) {
|
|
||||||
Select-String -Path $wfpDump -Pattern 'Nebula' -Context 0,80 -ErrorAction SilentlyContinue | Out-Host
|
|
||||||
}
|
|
||||||
throw
|
|
||||||
}
|
|
||||||
finally {
|
|
||||||
if (-not $lhProc.HasExited) {
|
|
||||||
Stop-Process -Id $lhProc.Id -Force -ErrorAction SilentlyContinue
|
|
||||||
$lhProc.WaitForExit(5000) | Out-Null
|
|
||||||
}
|
|
||||||
wsl -d $Distro -u root -- bash -c "pkill -f $WslDir/nebula 2>/dev/null; true" | Out-Null
|
|
||||||
# pkill returns 1 when no match and wsl propagates that; the smoke is done
|
|
||||||
# so we don't want it to leak into the script's exit code.
|
|
||||||
$global:LASTEXITCODE = 0
|
|
||||||
if ($peerProc -and -not $peerProc.HasExited) {
|
|
||||||
Stop-Process -Id $peerProc.Id -Force -ErrorAction SilentlyContinue
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -47,19 +47,6 @@ 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
|
||||||
@@ -93,28 +80,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 2000 &
|
docker exec host2 ncat -nklv 0.0.0.0 2000 &
|
||||||
docker exec host3 ncat -nklv 2000 &
|
docker exec host3 ncat -nklv 0.0.0.0 2000 &
|
||||||
docker exec host4 ncat -e '/usr/bin/echo helloagainfromhost4' -nkluv 4000 &
|
docker exec host4 ncat -nkluv 0.0.0.0 4000 &
|
||||||
docker exec host2 ncat -e '/usr/bin/echo host2' -nkluv 3000 &
|
docker exec host2 ncat -e '/usr/bin/echo host2' -nkluv 0.0.0.0 3000 &
|
||||||
docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 3000 &
|
docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 0.0.0.0 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 $HOST2_NIP
|
docker exec lighthouse1 ping -c1 192.168.100.2
|
||||||
docker exec lighthouse1 ping -c1 $HOST3_NIP
|
docker exec lighthouse1 ping -c1 192.168.100.3
|
||||||
|
|
||||||
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 $LIGHTHOUSE_NIP
|
docker exec host2 ping -c1 192.168.100.1
|
||||||
# Should fail because not allowed by host3 inbound firewall
|
# Should fail because not allowed by host3 inbound firewall
|
||||||
! docker exec host2 ping -c1 $HOST3_NIP -w5 || exit 1
|
! docker exec host2 ping -c1 192.168.100.3 -w5 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
@@ -122,34 +109,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 $HOST3_NIP 2000 || exit 1
|
! docker exec host2 ncat -nzv -w5 192.168.100.3 2000 || exit 1
|
||||||
! docker exec host2 ncat -nzuv -w5 $HOST3_NIP 3000 | grep -q host3 || exit 1
|
! docker exec host2 ncat -nzuv -w5 192.168.100.3 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 $LIGHTHOUSE_NIP
|
docker exec host3 ping -c1 192.168.100.1
|
||||||
docker exec host3 ping -c1 $HOST2_NIP
|
docker exec host3 ping -c1 192.168.100.2
|
||||||
|
|
||||||
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 $HOST2_NIP 2000
|
docker exec host3 ncat -nzv -w5 192.168.100.2 2000
|
||||||
docker exec host3 ncat -nzuv -w5 $HOST2_NIP 3000 | grep -q host2
|
docker exec host3 ncat -nzuv -w5 192.168.100.2 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 $LIGHTHOUSE_NIP
|
docker exec host4 ping -c1 192.168.100.1
|
||||||
# Should fail because not allowed by host4 outbound firewall
|
# Should fail because not allowed by host4 outbound firewall
|
||||||
! docker exec host4 ping -c1 $HOST2_NIP -w5 || exit 1
|
! docker exec host4 ping -c1 192.168.100.2 -w5 || exit 1
|
||||||
! docker exec host4 ping -c1 $HOST3_NIP -w5 || exit 1
|
! docker exec host4 ping -c1 192.168.100.3 -w5 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
@@ -157,10 +144,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 $HOST2_NIP 2000 || exit 1
|
! docker exec host4 ncat -nzv -w5 192.168.100.2 2000 || exit 1
|
||||||
! docker exec host4 ncat -nzv -w5 $HOST3_NIP 2000 || exit 1
|
! docker exec host4 ncat -nzv -w5 192.168.100.3 2000 || 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.2 3000 | grep -q host2 || exit 1
|
||||||
! docker exec host4 ncat -nzuv -w5 $HOST3_NIP 3000 | grep -q host3 || exit 1
|
! docker exec host4 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
@@ -168,11 +155,11 @@ echo " *** Testing conntrack"
|
|||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
|
|
||||||
# host4's outbound firewall only allows ICMP to the lighthouse, so host4
|
# host2 speaking to host4 on UDP 4000 should allow it to reply, when firewall rules would normally not permit this
|
||||||
# cannot initiate UDP to host2. Once host2 initiates a flow to host4:4000,
|
docker exec host2 sh -c "/usr/bin/echo host2 | ncat -nuv 192.168.100.4 4000"
|
||||||
# conntrack must let host4's listener reply on that flow. If it doesn't,
|
docker exec host2 ncat -e '/usr/bin/echo helloagainfromhost2' -nkluv 0.0.0.0 4000 &
|
||||||
# the echo back from host4 never reaches host2.
|
sleep 1
|
||||||
docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv $HOST4_NIP 4000" | grep -q helloagainfromhost4
|
docker exec host4 sh -c "/usr/bin/echo host4 | ncat -nuv 192.168.100.2 4000"
|
||||||
|
|
||||||
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 = "DefinedNet/netbsd10"
|
config.vm.box = "generic/netbsd9"
|
||||||
|
|
||||||
config.vm.synced_folder "../build", "/nebula", type: "rsync"
|
config.vm.synced_folder "../build", "/nebula", type: "rsync"
|
||||||
end
|
end
|
||||||
|
|||||||
+79
-102
@@ -13,28 +13,20 @@ on:
|
|||||||
- 'go.sum'
|
- 'go.sum'
|
||||||
jobs:
|
jobs:
|
||||||
|
|
||||||
static:
|
test-linux:
|
||||||
name: Static checks
|
name: Build all and test on ubuntu-linux
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.26'
|
go-version: '1.25'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Install goimports
|
- name: Build
|
||||||
run: go install golang.org/x/tools/cmd/goimports@latest
|
run: make all
|
||||||
|
|
||||||
- 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
|
||||||
@@ -42,114 +34,99 @@ jobs:
|
|||||||
- name: golangci-lint
|
- name: golangci-lint
|
||||||
uses: golangci/golangci-lint-action@v9
|
uses: golangci/golangci-lint-action@v9
|
||||||
with:
|
with:
|
||||||
version: v2.12
|
version: v2.5
|
||||||
|
|
||||||
test:
|
- name: Test
|
||||||
name: Test ${{ matrix.name }}
|
run: make test
|
||||||
runs-on: ${{ matrix.os }}
|
|
||||||
strategy:
|
- name: End 2 end
|
||||||
fail-fast: false
|
run: make e2evv
|
||||||
matrix:
|
|
||||||
include:
|
- name: Build test mobile
|
||||||
- name: linux
|
run: make build-test-mobile
|
||||||
os: ubuntu-latest
|
|
||||||
build-cmd: go build ./cmd/nebula ./cmd/nebula-cert
|
- uses: actions/upload-artifact@v6
|
||||||
test-cmd: make test
|
with:
|
||||||
e2e-cmd: make e2evv
|
name: e2e packet flow linux-latest
|
||||||
- name: linux-boringcrypto
|
path: e2e/mermaid/linux-latest
|
||||||
os: ubuntu-latest
|
if-no-files-found: warn
|
||||||
build-cmd: make boringcrypto
|
|
||||||
test-cmd: make boringcrypto test
|
test-linux-boringcrypto:
|
||||||
e2e-cmd: make boringcrypto e2evv
|
name: Build and test on linux with boringcrypto
|
||||||
- name: linux-fips140
|
runs-on: ubuntu-latest
|
||||||
os: ubuntu-latest
|
|
||||||
build-cmd: make fips140-all
|
|
||||||
test-cmd: make fips140-all GOALS=test
|
|
||||||
e2e-cmd: make fips140-all GOALS=e2evv
|
|
||||||
- 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@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.26'
|
go-version: '1.25'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
run: ${{ matrix.build-cmd }}
|
run: make bin-boringcrypto
|
||||||
|
|
||||||
- name: Cross-build darwin-amd64
|
|
||||||
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: Test
|
- name: Test
|
||||||
run: ${{ matrix.test-cmd }}
|
run: make test-boringcrypto
|
||||||
|
|
||||||
- name: End 2 end
|
- name: End 2 end
|
||||||
if: matrix.e2e-cmd != ''
|
run: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0"
|
||||||
run: ${{ matrix.e2e-cmd }}
|
|
||||||
|
|
||||||
- uses: actions/upload-artifact@v7
|
test-linux-pkcs11:
|
||||||
if: matrix.e2e-cmd != '' && always()
|
name: Build and test on linux with pkcs11
|
||||||
with:
|
|
||||||
name: e2e packet flow ${{ matrix.name }}
|
|
||||||
path: e2e/mermaid/
|
|
||||||
if-no-files-found: warn
|
|
||||||
|
|
||||||
cross-build:
|
|
||||||
name: Cross-build ${{ matrix.name }}
|
|
||||||
runs-on: ubuntu-latest
|
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:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.26'
|
go-version: '1.25'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build ${{ matrix.name }}
|
- name: Build
|
||||||
run: make -j"$(nproc)" ${{ matrix.make-target }}
|
run: make bin-pkcs11
|
||||||
|
|
||||||
finish:
|
- name: Test
|
||||||
name: CI status
|
run: make test-pkcs11
|
||||||
if: always()
|
|
||||||
needs: [static, test, cross-build]
|
test:
|
||||||
runs-on: ubuntu-latest
|
name: Build and test on ${{ matrix.os }}
|
||||||
|
runs-on: ${{ matrix.os }}
|
||||||
|
strategy:
|
||||||
|
matrix:
|
||||||
|
os: [windows-latest, macos-latest]
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- name: Fail if any upstream job failed
|
- uses: actions/checkout@v6
|
||||||
if: contains(needs.*.result, 'failure') || contains(needs.*.result, 'cancelled')
|
|
||||||
run: |
|
|
||||||
echo "upstream results: ${{ toJSON(needs) }}"
|
|
||||||
exit 1
|
|
||||||
|
|
||||||
- name: All upstream jobs passed
|
- uses: actions/setup-go@v6
|
||||||
run: echo "ok"
|
with:
|
||||||
|
go-version: '1.25'
|
||||||
|
check-latest: true
|
||||||
|
|
||||||
|
- name: Build nebula
|
||||||
|
run: go build ./cmd/nebula
|
||||||
|
|
||||||
|
- name: Build nebula-cert
|
||||||
|
run: go build ./cmd/nebula-cert
|
||||||
|
|
||||||
|
- name: Vet
|
||||||
|
run: make vet
|
||||||
|
|
||||||
|
- name: golangci-lint
|
||||||
|
uses: golangci/golangci-lint-action@v9
|
||||||
|
with:
|
||||||
|
version: v2.5
|
||||||
|
|
||||||
|
- name: Test
|
||||||
|
run: make test
|
||||||
|
|
||||||
|
- name: End 2 end
|
||||||
|
run: make e2evv
|
||||||
|
|
||||||
|
- uses: actions/upload-artifact@v6
|
||||||
|
with:
|
||||||
|
name: e2e packet flow ${{ matrix.os }}
|
||||||
|
path: e2e/mermaid/${{ matrix.os }}
|
||||||
|
if-no-files-found: warn
|
||||||
|
|||||||
@@ -2,21 +2,7 @@ version: "2"
|
|||||||
linters:
|
linters:
|
||||||
default: none
|
default: none
|
||||||
enable:
|
enable:
|
||||||
- sloglint
|
|
||||||
- testifylint
|
- testifylint
|
||||||
settings:
|
|
||||||
sloglint:
|
|
||||||
# Enforce key-value pair form for Info/Debug/Warn/Error/Log/With and
|
|
||||||
# the package-level slog equivalents. Use l.Log(ctx, level, ...) for
|
|
||||||
# custom levels instead of LogAttrs when you can.
|
|
||||||
#
|
|
||||||
# LogAttrs is also flagged by this rule because it takes ...slog.Attr;
|
|
||||||
# the few legitimate sites (where attrs is built up as a []slog.Attr)
|
|
||||||
# carry a //nolint:sloglint with rationale.
|
|
||||||
kv-only: true
|
|
||||||
# no-mixed-args is on by default: forbids mixing kv and attrs in one call.
|
|
||||||
# discard-handler is on by default (since Go 1.24): suggests
|
|
||||||
# slog.DiscardHandler over slog.NewTextHandler(io.Discard, nil).
|
|
||||||
exclusions:
|
exclusions:
|
||||||
generated: lax
|
generated: lax
|
||||||
presets:
|
presets:
|
||||||
|
|||||||
@@ -60,29 +60,6 @@ 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
|
|
||||||
|
|
||||||
# Based on section 2.2 of the Go Cryptographic Module CVMP Security Policy #5247
|
|
||||||
ALL_FIPS140 = linux-amd64-fips140 \
|
|
||||||
linux-arm64-fips140 \
|
|
||||||
windows-amd64-fips140 \
|
|
||||||
windows-arm64-fips140 \
|
|
||||||
darwin-arm64-fips140 \
|
|
||||||
freebsd-amd64-fips140 \
|
|
||||||
linux-arm-7-fips140 \
|
|
||||||
linux-mips64-fips140 \
|
|
||||||
linux-ppc64le-fips140
|
|
||||||
|
|
||||||
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
|
||||||
|
|
||||||
@@ -105,35 +82,6 @@ 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)
|
||||||
@@ -148,8 +96,6 @@ release-netbsd: $(ALL_NETBSD:%=build/nebula-%.tar.gz)
|
|||||||
|
|
||||||
release-boringcrypto: build/nebula-linux-$(shell go env GOARCH)-boringcrypto.tar.gz
|
release-boringcrypto: build/nebula-linux-$(shell go env GOARCH)-boringcrypto.tar.gz
|
||||||
|
|
||||||
release-fips140: $(ALL_FIPS140:%=build/nebula-%.tar.gz)
|
|
||||||
|
|
||||||
BUILD_ARGS += -trimpath
|
BUILD_ARGS += -trimpath
|
||||||
|
|
||||||
bin-windows: build/windows-amd64/nebula.exe build/windows-amd64/nebula-cert.exe
|
bin-windows: build/windows-amd64/nebula.exe build/windows-amd64/nebula-cert.exe
|
||||||
@@ -170,20 +116,17 @@ bin-freebsd-arm64: build/freebsd-arm64/nebula build/freebsd-arm64/nebula-cert
|
|||||||
bin-boringcrypto: build/linux-$(shell go env GOARCH)-boringcrypto/nebula build/linux-$(shell go env GOARCH)-boringcrypto/nebula-cert
|
bin-boringcrypto: build/linux-$(shell go env GOARCH)-boringcrypto/nebula build/linux-$(shell go env GOARCH)-boringcrypto/nebula-cert
|
||||||
mv $? .
|
mv $? .
|
||||||
|
|
||||||
bin-fips140: build/linux-$(shell go env GOARCH)-fips140/nebula build/linux-$(shell go env GOARCH)-fips140/nebula-cert
|
|
||||||
mv $? .
|
|
||||||
|
|
||||||
bin-pkcs11: BUILD_ARGS += -tags pkcs11
|
bin-pkcs11: BUILD_ARGS += -tags pkcs11
|
||||||
bin-pkcs11: CGO_ENABLED = 1
|
bin-pkcs11: CGO_ENABLED = 1
|
||||||
bin-pkcs11: bin
|
bin-pkcs11: bin
|
||||||
|
|
||||||
bin:
|
bin:
|
||||||
$(GOENV) go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
||||||
$(GOENV) go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
||||||
|
|
||||||
install:
|
install:
|
||||||
$(GOENV) go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ${NEBULA_CMD_PATH}
|
go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ${NEBULA_CMD_PATH}
|
||||||
$(GOENV) go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ./cmd/nebula-cert
|
go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ./cmd/nebula-cert
|
||||||
|
|
||||||
build/linux-arm-%: GOENV += GOARM=$(word 3, $(subst -, ,$*))
|
build/linux-arm-%: GOENV += GOARM=$(word 3, $(subst -, ,$*))
|
||||||
build/linux-mips-%: GOENV += GOMIPS=$(word 3, $(subst -, ,$*))
|
build/linux-mips-%: GOENV += GOMIPS=$(word 3, $(subst -, ,$*))
|
||||||
@@ -194,11 +137,8 @@ build/linux-mips-softfloat/%: LDFLAGS += -s -w
|
|||||||
# boringcrypto
|
# boringcrypto
|
||||||
build/linux-amd64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
build/linux-amd64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
||||||
build/linux-arm64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
build/linux-arm64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
||||||
|
build/linux-amd64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
||||||
# fips140
|
build/linux-arm64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
||||||
FIPSVERSION = v1.0.0
|
|
||||||
$(foreach _rule, $(ALL_FIPS140), build/$(_rule)/%): GOENV += GOFIPS140=$(FIPSVERSION)
|
|
||||||
$(foreach _rule, $(ALL_FIPS140), build/$(_rule)/%): BUILD_ARGS += -tags fips140enforce
|
|
||||||
|
|
||||||
build/%/nebula: .FORCE
|
build/%/nebula: .FORCE
|
||||||
GOOS=$(firstword $(subst -, , $*)) \
|
GOOS=$(firstword $(subst -, , $*)) \
|
||||||
@@ -229,7 +169,10 @@ vet:
|
|||||||
go vet $(VET_FLAGS) -v ./...
|
go vet $(VET_FLAGS) -v ./...
|
||||||
|
|
||||||
test:
|
test:
|
||||||
$(TEST_ENV) go test $(TEST_FLAGS) -v ./...
|
go test -v ./...
|
||||||
|
|
||||||
|
test-boringcrypto:
|
||||||
|
GOEXPERIMENT=boringcrypto CGO_ENABLED=1 go test -ldflags "-checklinkname=0" -v ./...
|
||||||
|
|
||||||
test-pkcs11:
|
test-pkcs11:
|
||||||
CGO_ENABLED=1 go test -v -tags pkcs11 ./...
|
CGO_ENABLED=1 go test -v -tags pkcs11 ./...
|
||||||
@@ -272,72 +215,26 @@ ifeq ($(words $(MAKECMDGOALS)),1)
|
|||||||
@$(MAKE) service ${.DEFAULT_GOAL} --no-print-directory
|
@$(MAKE) service ${.DEFAULT_GOAL} --no-print-directory
|
||||||
endif
|
endif
|
||||||
|
|
||||||
# Useful to chain together, like:
|
|
||||||
# - make fips140 e2evv
|
|
||||||
# - make fips140 smoke-docker
|
|
||||||
# Use `release-fips140` to build release binaries
|
|
||||||
fips140:
|
|
||||||
@echo > $(NULL_FILE)
|
|
||||||
ifeq ($(strip $(GOFIPS140)),)
|
|
||||||
$(eval GOFIPS140 = $(FIPSVERSION))
|
|
||||||
endif
|
|
||||||
$(eval GOENV += GOFIPS140=$(GOFIPS140))
|
|
||||||
$(eval BUILD_ARGS += -tags fips140enforce)
|
|
||||||
$(eval TEST_ENV += $(GOENV))
|
|
||||||
$(eval CURVE = P256)
|
|
||||||
ifeq ($(words $(MAKECMDGOALS)),1)
|
|
||||||
@$(MAKE) fips140 GOFIPS140=$(GOFIPS140) ${.DEFAULT_GOAL} --no-print-directory
|
|
||||||
endif
|
|
||||||
|
|
||||||
# To test the future pending module, use like `make fips140-latest test`
|
|
||||||
ALL_GOFIPS140 = v1.0.0 v1.26.0 latest
|
|
||||||
define FIPS140_rule
|
|
||||||
fips140-$(1): GOFIPS140 = $(1)
|
|
||||||
fips140-$(1): fips140
|
|
||||||
endef
|
|
||||||
$(foreach _rule, $(ALL_GOFIPS140), $(eval $(call FIPS140_rule,$(_rule))))
|
|
||||||
|
|
||||||
# Iterate and run the goals for all fips versions, like `make fips140-all GOALS=test`
|
|
||||||
fips140-all:
|
|
||||||
@$(foreach _v,$(ALL_GOFIPS140),$(MAKE) fips140-$(_v) $(GOALS) &&) true
|
|
||||||
|
|
||||||
# Useful to chain together, like:
|
|
||||||
# - make boringcrypto e2evv
|
|
||||||
# - make boringcrypto smoke-docker
|
|
||||||
# Use `release-boringcrypto` or `bin-boringcrypto` to build release binaries
|
|
||||||
boringcrypto:
|
|
||||||
@echo > $(NULL_FILE)
|
|
||||||
$(eval GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1)
|
|
||||||
$(eval TEST_ENV += $(GOENV))
|
|
||||||
$(eval CURVE = P256)
|
|
||||||
ifeq ($(words $(MAKECMDGOALS)),1)
|
|
||||||
@$(MAKE) boringcrypto ${.DEFAULT_GOAL} --no-print-directory
|
|
||||||
endif
|
|
||||||
|
|
||||||
bin-docker: bin build/linux-amd64/nebula build/linux-amd64/nebula-cert
|
bin-docker: bin build/linux-amd64/nebula build/linux-amd64/nebula-cert
|
||||||
|
|
||||||
smoke-docker: BUILD_ARGS += -race
|
|
||||||
smoke-docker: GOENV += CGO_ENABLED=1
|
|
||||||
smoke-docker: bin-docker
|
smoke-docker: bin-docker
|
||||||
# This is so we can limit `fips140` smoke test to just P256 curve.
|
cd .github/workflows/smoke/ && ./build.sh
|
||||||
if [ "$(CURVE)" != "P256" ]; then cd .github/workflows/smoke/ && $(GOENV) ./build.sh; fi
|
cd .github/workflows/smoke/ && ./smoke.sh
|
||||||
if [ "$(CURVE)" != "P256" ]; then cd .github/workflows/smoke/ && $(GOENV) ./smoke.sh; fi
|
cd .github/workflows/smoke/ && NAME="smoke-p256" CURVE="P256" ./build.sh
|
||||||
cd .github/workflows/smoke/ && $(GOENV) NAME="smoke-p256" CURVE="P256" ./build.sh
|
cd .github/workflows/smoke/ && NAME="smoke-p256" ./smoke.sh
|
||||||
cd .github/workflows/smoke/ && $(GOENV) NAME="smoke-p256" ./smoke.sh
|
|
||||||
|
|
||||||
smoke-relay-docker: BUILD_ARGS += -race
|
|
||||||
smoke-relay-docker: GOENV += CGO_ENABLED=1
|
|
||||||
smoke-relay-docker: bin-docker
|
smoke-relay-docker: bin-docker
|
||||||
cd .github/workflows/smoke/ && $(GOENV) ./build-relay.sh
|
cd .github/workflows/smoke/ && ./build-relay.sh
|
||||||
cd .github/workflows/smoke/ && $(GOENV) ./smoke-relay.sh
|
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
||||||
|
|
||||||
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
|
smoke-docker-race: BUILD_ARGS = -race
|
||||||
smoke-docker-ipv6: smoke-docker
|
smoke-docker-race: CGO_ENABLED = 1
|
||||||
|
smoke-docker-race: smoke-docker
|
||||||
|
|
||||||
smoke-vagrant/%: bin-docker build/%/nebula
|
smoke-vagrant/%: bin-docker build/%/nebula
|
||||||
cd .github/workflows/smoke/ && ./build.sh $*
|
cd .github/workflows/smoke/ && ./build.sh $*
|
||||||
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
||||||
|
|
||||||
.FORCE:
|
.FORCE:
|
||||||
.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 bin-windows bin-windows-arm64 bin-darwin bin-freebsd bin-freebsd-arm64 bin-boringcrypto bin-fips140 bin-pkcs11 bin-docker boringcrypto build-test-mobile docker e2e e2ev e2evv e2evvv e2evvvv e2e-bench fips140 fips140-all $(ALL_GOFIPS140:%=fips140-%) install proto release release-linux release-freebsd release-openbsd release-netbsd release-boringcrypto release-fips140 service smoke-docker smoke-relay-docker smoke-docker-ipv6 test test-pkcs11 test-cov-html vet smoke-vagrant/%
|
.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/%
|
||||||
.DEFAULT_GOAL := bin
|
.DEFAULT_GOAL := bin
|
||||||
|
|||||||
@@ -145,27 +145,17 @@ To build nebula for a specific platform (ex, Windows):
|
|||||||
|
|
||||||
See the [Makefile](Makefile) for more details on build targets
|
See the [Makefile](Makefile) for more details on build targets
|
||||||
|
|
||||||
## Curve P256 and FIPS 140-3 mode
|
## Curve P256 and BoringCrypto
|
||||||
|
|
||||||
The default curve used for cryptographic handshakes and signatures is Curve25519. This is the recommended setting for most users. If your deployment has certain compliance requirements, you have the option of creating your CA using `nebula-cert ca -curve P256` to use NIST Curve P256. The CA will then sign certificates using ECDSA P256, and any hosts using these certificates will use P256 for ECDH handshakes.
|
The default curve used for cryptographic handshakes and signatures is Curve25519. This is the recommended setting for most users. If your deployment has certain compliance requirements, you have the option of creating your CA using `nebula-cert ca -curve P256` to use NIST Curve P256. The CA will then sign certificates using ECDSA P256, and any hosts using these certificates will use P256 for ECDH handshakes.
|
||||||
|
|
||||||
Nebula can be built to support the [FIPS 140-3](https://go.dev/doc/security/fips140) mode of Go by running either of the following make targets. (This sets GOFIPS140=v1.0.0, which must be done at compile time so that the correct AES-GCM can be used for FIPS 140-3 enforcement mode).
|
In addition, Nebula can be built using the [BoringCrypto GOEXPERIMENT](https://github.com/golang/go/blob/go1.20/src/crypto/internal/boring/README.md) by running either of the following make targets:
|
||||||
|
|
||||||
```sh
|
|
||||||
make fips140
|
|
||||||
make fips140 test
|
|
||||||
make release-fips140
|
|
||||||
```
|
|
||||||
|
|
||||||
Nebula can also be built using the [BoringCrypto GOEXPERIMENT](https://github.com/golang/go/blob/go1.20/src/crypto/internal/boring/README.md) by running either of the following make targets.
|
|
||||||
|
|
||||||
```sh
|
```sh
|
||||||
make bin-boringcrypto
|
make bin-boringcrypto
|
||||||
make release-boringcrypto
|
make release-boringcrypto
|
||||||
```
|
```
|
||||||
|
|
||||||
NOTE: boringcrypto support is deprecated and will be removed in the next release. Users should migrate to the native FIPS 140-3 mode described above.
|
|
||||||
|
|
||||||
This is not the recommended default deployment, but may be useful based on your compliance requirements.
|
This is not the recommended default deployment, but may be useful based on your compliance requirements.
|
||||||
|
|
||||||
## Credits
|
## Credits
|
||||||
|
|||||||
@@ -1,43 +1,23 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"fmt"
|
|
||||||
"log/slog"
|
|
||||||
"math"
|
|
||||||
mathbits "math/bits"
|
|
||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
)
|
)
|
||||||
|
|
||||||
const bitsPerWord = 64
|
|
||||||
|
|
||||||
// Bits is a sliding-window anti-replay tracker. The window is stored as a
|
|
||||||
// circular bitmap packed into uint64 words (8x denser than a []bool), so a
|
|
||||||
// length-N window costs N/8 bytes. length must be a power of two.
|
|
||||||
type Bits struct {
|
type Bits struct {
|
||||||
length uint64
|
length uint64
|
||||||
lengthMask uint64
|
|
||||||
current uint64
|
current uint64
|
||||||
bits []uint64
|
bits []bool
|
||||||
lostCounter metrics.Counter
|
lostCounter metrics.Counter
|
||||||
dupeCounter metrics.Counter
|
dupeCounter metrics.Counter
|
||||||
outOfWindowCounter metrics.Counter
|
outOfWindowCounter metrics.Counter
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewBits(length uint64) *Bits {
|
func NewBits(bits uint64) *Bits {
|
||||||
if length == 0 || length&(length-1) != 0 {
|
|
||||||
panic(fmt.Sprintf("Bits length must be a power of two, got %d", length))
|
|
||||||
}
|
|
||||||
|
|
||||||
nWords := length / bitsPerWord
|
|
||||||
if nWords == 0 {
|
|
||||||
nWords = 1
|
|
||||||
}
|
|
||||||
b := &Bits{
|
b := &Bits{
|
||||||
length: length,
|
length: bits,
|
||||||
lengthMask: length - 1,
|
bits: make([]bool, bits, bits),
|
||||||
bits: make([]uint64, nWords),
|
|
||||||
current: 0,
|
current: 0,
|
||||||
lostCounter: metrics.GetOrRegisterCounter("network.packets.lost", nil),
|
lostCounter: metrics.GetOrRegisterCounter("network.packets.lost", nil),
|
||||||
dupeCounter: metrics.GetOrRegisterCounter("network.packets.duplicate", nil),
|
dupeCounter: metrics.GetOrRegisterCounter("network.packets.duplicate", nil),
|
||||||
@@ -45,219 +25,88 @@ func NewBits(length uint64) *Bits {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// There is no counter value 0, mark it to avoid counting a lost packet later.
|
// There is no counter value 0, mark it to avoid counting a lost packet later.
|
||||||
b.bits[0] = 1
|
b.bits[0] = true
|
||||||
|
b.current = 0
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *Bits) get(i uint64) bool {
|
func (b *Bits) Check(l *logrus.Logger, i uint64) bool {
|
||||||
pos := i & b.lengthMask
|
|
||||||
//bit-shifting by 6 because i is a bit index, not a u64 index, and we need to find the u64 without bit in it
|
|
||||||
return b.bits[pos>>6]&(uint64(1)<<(pos&63)) != 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (b *Bits) set(i uint64) {
|
|
||||||
pos := i & b.lengthMask
|
|
||||||
b.bits[pos>>6] |= uint64(1) << (pos & 63)
|
|
||||||
}
|
|
||||||
|
|
||||||
// clearRange clears `count` bits starting at circular position `startPos`
|
|
||||||
// (already masked to [0, length)) and returns how many of them were set
|
|
||||||
// before the clear. count must be in [1, length].
|
|
||||||
func (b *Bits) clearRange(startPos, count uint64) uint64 {
|
|
||||||
wasSet := uint64(0)
|
|
||||||
if count >= b.length {
|
|
||||||
for _, w := range b.bits {
|
|
||||||
wasSet += uint64(mathbits.OnesCount64(w))
|
|
||||||
}
|
|
||||||
clear(b.bits)
|
|
||||||
return wasSet
|
|
||||||
}
|
|
||||||
|
|
||||||
pos := startPos
|
|
||||||
remaining := count
|
|
||||||
|
|
||||||
// handle the potential partial word before pos becomes u64 aligned
|
|
||||||
word := pos >> 6
|
|
||||||
bit := pos & 63
|
|
||||||
take := uint64(64) - bit
|
|
||||||
if take > remaining {
|
|
||||||
take = remaining
|
|
||||||
}
|
|
||||||
if take > b.length-pos {
|
|
||||||
take = b.length - pos
|
|
||||||
}
|
|
||||||
var mask uint64
|
|
||||||
if take == 64 {
|
|
||||||
mask = math.MaxUint64
|
|
||||||
} else {
|
|
||||||
mask = ((uint64(1) << take) - 1) << bit
|
|
||||||
}
|
|
||||||
wasSet += uint64(mathbits.OnesCount64(b.bits[word] & mask))
|
|
||||||
b.bits[word] &^= mask
|
|
||||||
remaining -= take
|
|
||||||
pos = (pos + take) & b.lengthMask
|
|
||||||
|
|
||||||
// Clear whole words, keeping track of the number of set bits
|
|
||||||
for remaining >= 64 {
|
|
||||||
word = pos >> 6
|
|
||||||
wasSet += uint64(mathbits.OnesCount64(b.bits[word]))
|
|
||||||
b.bits[word] = 0
|
|
||||||
remaining -= 64
|
|
||||||
pos = (pos + 64) & b.lengthMask
|
|
||||||
}
|
|
||||||
|
|
||||||
// Clear the remaining partial word
|
|
||||||
if remaining > 0 {
|
|
||||||
word = pos >> 6
|
|
||||||
mask = (uint64(1) << remaining) - 1
|
|
||||||
wasSet += uint64(mathbits.OnesCount64(b.bits[word] & mask))
|
|
||||||
b.bits[word] &^= mask
|
|
||||||
}
|
|
||||||
|
|
||||||
return wasSet
|
|
||||||
}
|
|
||||||
|
|
||||||
func (b *Bits) strictlyWithinWindow(i uint64) bool {
|
|
||||||
// Handle the case where the window hasn't slid yet. This avoids u64 underflow.
|
|
||||||
inWarmup := b.current < b.length
|
|
||||||
if i < b.length && inWarmup {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// Next, if the packet is in-window, see if we've seen it before
|
|
||||||
if i > b.current-b.length {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
return false //not within window!
|
|
||||||
}
|
|
||||||
|
|
||||||
// Check returns true if i is within (or way out in front of) the window, and not a replay
|
|
||||||
func (b *Bits) Check(l *slog.Logger, i uint64) bool {
|
|
||||||
// If i is the next number, return true.
|
// If i is the next number, return true.
|
||||||
if i > b.current {
|
if i > b.current {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
if b.strictlyWithinWindow(i) {
|
// If i is within the window, check if it's been set already.
|
||||||
return !b.get(i)
|
if i > b.current-b.length || i < b.length && b.current < b.length {
|
||||||
|
return !b.bits[i%b.length]
|
||||||
}
|
}
|
||||||
|
|
||||||
// Not within the window
|
// Not within the window
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
if l.Level >= logrus.DebugLevel {
|
||||||
l.Debug("rejected a packet (top)", "current", b.current, "incoming", i)
|
l.Debugf("rejected a packet (top) %d %d\n", b.current, i)
|
||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// Update has three branches:
|
func (b *Bits) Update(l *logrus.Logger, i uint64) bool {
|
||||||
// - i == b.current+1: fast path; advance the cursor by one and lose-count
|
// If i is the next number, return true and update current.
|
||||||
// the slot we just stomped (only past warmup; see the i > b.length guard
|
|
||||||
// below).
|
|
||||||
// - i > b.current+1: jump path; clear all slots between current and i
|
|
||||||
// (or up to a full window's worth, whichever is smaller) via clearRange,
|
|
||||||
// then mark i. Two arms here: a warmup arm that handles the very first
|
|
||||||
// window before the cursor has slid, and a steady-state arm that treats
|
|
||||||
// every cleared empty slot as a lost packet.
|
|
||||||
// - i <= b.current: in-window check for duplicates; out-of-window otherwise.
|
|
||||||
//
|
|
||||||
// NewBits seeds bits[0]=1 so counter 0 looks "received" — Update never
|
|
||||||
// clears that marker during warmup (clearRange skips position 0 when
|
|
||||||
// startPos=1), and once b.current >= b.length the marker is no longer
|
|
||||||
// consulted. The marker prevents a fictitious "lost" hit on the first real
|
|
||||||
// counter.
|
|
||||||
func (b *Bits) Update(l *slog.Logger, i uint64) bool {
|
|
||||||
// Fast path: i is the next expected counter. Split out so the function
|
|
||||||
// stays small and avoids paying for the slow paths' slog argument-build
|
|
||||||
// stack frame on every call. The bit read/test/write is inlined to
|
|
||||||
// touch the backing word once.
|
|
||||||
if i == b.current+1 {
|
if i == b.current+1 {
|
||||||
pos := i & b.lengthMask
|
// Check if the oldest bit was lost since we are shifting the window by 1 and occupying it with this counter
|
||||||
word := pos >> 6
|
// The very first window can only be tracked as lost once we are on the 2nd window or greater
|
||||||
mask := uint64(1) << (pos & 63)
|
if b.bits[i%b.length] == false && i > b.length {
|
||||||
w := b.bits[word]
|
|
||||||
if i > b.length && w&mask == 0 {
|
|
||||||
b.lostCounter.Inc(1)
|
b.lostCounter.Inc(1)
|
||||||
}
|
}
|
||||||
b.bits[word] = w | mask
|
b.bits[i%b.length] = true
|
||||||
b.current = i
|
b.current = i
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
return b.updateSlow(l, i)
|
|
||||||
}
|
|
||||||
|
|
||||||
// updateSlow handles jumps, in-window backfill, dupes, and out-of-window.
|
|
||||||
func (b *Bits) updateSlow(l *slog.Logger, i uint64) bool {
|
|
||||||
// If i is a jump, adjust the window, record lost, update current, and return true
|
// If i is a jump, adjust the window, record lost, update current, and return true
|
||||||
if i > b.current {
|
if i > b.current {
|
||||||
end := i
|
lost := int64(0)
|
||||||
if end > b.current+b.length {
|
// Zero out the bits between the current and the new counter value, limited by the window size,
|
||||||
end = b.current + b.length
|
// since the window is shifting
|
||||||
}
|
for n := b.current + 1; n <= min(i, b.current+b.length); n++ {
|
||||||
count := end - b.current
|
if b.bits[n%b.length] == false && n > b.length {
|
||||||
startPos := (b.current + 1) & b.lengthMask
|
lost++
|
||||||
|
|
||||||
var lost int64
|
|
||||||
if b.current >= b.length {
|
|
||||||
// Steady state: every cleared slot is past warmup, so any unset
|
|
||||||
// bit we evict is a lost packet from the previous cycle.
|
|
||||||
wasSet := b.clearRange(startPos, count)
|
|
||||||
lost = int64(count) - int64(wasSet)
|
|
||||||
} else {
|
|
||||||
// Warmup (the very first window). Some cleared slots represent
|
|
||||||
// packets <= length where eviction is not "lost" in the usual
|
|
||||||
// sense. This branch is taken at most once per connection so we
|
|
||||||
// don't bother optimizing it.
|
|
||||||
for n := b.current + 1; n <= end; n++ {
|
|
||||||
if !b.get(n) && n > b.length {
|
|
||||||
lost++
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
b.clearRange(startPos, count)
|
b.bits[n%b.length] = false
|
||||||
}
|
}
|
||||||
|
|
||||||
// Anything past the new window can never be backfilled, so it's lost.
|
// Only record any skipped packets as a result of the window moving further than the window length
|
||||||
if i > b.current+b.length {
|
// Any loss within the new window will be accounted for in future calls
|
||||||
lost += int64(i - b.current - b.length)
|
lost += max(0, int64(i-b.current-b.length))
|
||||||
}
|
|
||||||
b.lostCounter.Inc(lost)
|
b.lostCounter.Inc(lost)
|
||||||
|
|
||||||
b.set(i)
|
b.bits[i%b.length] = true
|
||||||
b.current = i
|
b.current = i
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// If i is within the current window but below the current counter, check to see if it's a duplicate
|
// If i is within the current window but below the current counter,
|
||||||
if b.strictlyWithinWindow(i) {
|
// Check to see if it's a duplicate
|
||||||
pos := i & b.lengthMask
|
if i > b.current-b.length || i < b.length && b.current < b.length {
|
||||||
word := pos >> 6
|
if b.current == i || b.bits[i%b.length] == true {
|
||||||
mask := uint64(1) << (pos & 63)
|
if l.Level >= logrus.DebugLevel {
|
||||||
w := b.bits[word]
|
l.WithField("receiveWindow", m{"accepted": false, "currentCounter": b.current, "incomingCounter": i, "reason": "duplicate"}).
|
||||||
if b.current == i || w&mask != 0 {
|
Debug("Receive window")
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
l.Debug("Receive window",
|
|
||||||
"accepted", false,
|
|
||||||
"currentCounter", b.current,
|
|
||||||
"incomingCounter", i,
|
|
||||||
"reason", "duplicate",
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
b.dupeCounter.Inc(1)
|
b.dupeCounter.Inc(1)
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
b.bits[word] = w | mask
|
b.bits[i%b.length] = true
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// In all other cases, fail and don't change current.
|
// In all other cases, fail and don't change current.
|
||||||
b.outOfWindowCounter.Inc(1)
|
b.outOfWindowCounter.Inc(1)
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
if l.Level >= logrus.DebugLevel {
|
||||||
l.Debug("Receive window",
|
l.WithField("accepted", false).
|
||||||
"accepted", false,
|
WithField("currentCounter", b.current).
|
||||||
"currentCounter", b.current,
|
WithField("incomingCounter", i).
|
||||||
"incomingCounter", i,
|
WithField("reason", "nonsense").
|
||||||
"reason", "nonsense",
|
Debug("Receive window")
|
||||||
)
|
|
||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|||||||
+129
-276
@@ -7,79 +7,61 @@ import (
|
|||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
|
|
||||||
// snapshot returns the bitmap as a []bool of length b.length, for readable
|
|
||||||
// test assertions against the now-packed []uint64 storage.
|
|
||||||
func (b *Bits) snapshot() []bool {
|
|
||||||
out := make([]bool, b.length)
|
|
||||||
for i := uint64(0); i < b.length; i++ {
|
|
||||||
out[i] = b.get(i)
|
|
||||||
}
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBitsRequiresPowerOfTwo(t *testing.T) {
|
|
||||||
assert.Panics(t, func() { NewBits(10) })
|
|
||||||
assert.Panics(t, func() { NewBits(0) })
|
|
||||||
assert.NotPanics(t, func() { NewBits(1) })
|
|
||||||
assert.NotPanics(t, func() { NewBits(16) })
|
|
||||||
assert.NotPanics(t, func() { NewBits(1024) })
|
|
||||||
assert.NotPanics(t, func() { NewBits(16384) })
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBits(t *testing.T) {
|
func TestBits(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
b := NewBits(16)
|
b := NewBits(10)
|
||||||
assert.EqualValues(t, 16, b.length)
|
|
||||||
|
// make sure it is the right size
|
||||||
|
assert.Len(t, b.bits, 10)
|
||||||
|
|
||||||
// This is initialized to zero - receive one. This should work.
|
// This is initialized to zero - receive one. This should work.
|
||||||
assert.True(t, b.Check(l, 1))
|
assert.True(t, b.Check(l, 1))
|
||||||
assert.True(t, b.Update(l, 1))
|
assert.True(t, b.Update(l, 1))
|
||||||
assert.EqualValues(t, 1, b.current)
|
assert.EqualValues(t, 1, b.current)
|
||||||
g := []bool{true, true, false, false, false, false, false, false, false, false, false, false, false, false, false, false}
|
g := []bool{true, true, false, false, false, false, false, false, false, false}
|
||||||
assert.Equal(t, g, b.snapshot())
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Receive two
|
// Receive two
|
||||||
assert.True(t, b.Check(l, 2))
|
assert.True(t, b.Check(l, 2))
|
||||||
assert.True(t, b.Update(l, 2))
|
assert.True(t, b.Update(l, 2))
|
||||||
assert.EqualValues(t, 2, b.current)
|
assert.EqualValues(t, 2, b.current)
|
||||||
g = []bool{true, true, true, false, false, false, false, false, false, false, false, false, false, false, false, false}
|
g = []bool{true, true, true, false, false, false, false, false, false, false}
|
||||||
assert.Equal(t, g, b.snapshot())
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Receive two again - it will fail
|
// Receive two again - it will fail
|
||||||
assert.False(t, b.Check(l, 2))
|
assert.False(t, b.Check(l, 2))
|
||||||
assert.False(t, b.Update(l, 2))
|
assert.False(t, b.Update(l, 2))
|
||||||
assert.EqualValues(t, 2, b.current)
|
assert.EqualValues(t, 2, b.current)
|
||||||
|
|
||||||
// Jump ahead to 25, which clears the window and sets slot 25%16 = 9.
|
// Jump ahead to 15, which should clear everything and set the 6th element
|
||||||
assert.True(t, b.Check(l, 25))
|
assert.True(t, b.Check(l, 15))
|
||||||
assert.True(t, b.Update(l, 25))
|
assert.True(t, b.Update(l, 15))
|
||||||
assert.EqualValues(t, 25, b.current)
|
assert.EqualValues(t, 15, b.current)
|
||||||
g = []bool{false, false, false, false, false, false, false, false, false, true, false, false, false, false, false, false}
|
g = []bool{false, false, false, false, false, true, false, false, false, false}
|
||||||
assert.Equal(t, g, b.snapshot())
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Mark 24, which is in window (current 25, length 16, window covers [10,25]).
|
// Mark 14, which is allowed because it is in the window
|
||||||
assert.True(t, b.Check(l, 24))
|
assert.True(t, b.Check(l, 14))
|
||||||
assert.True(t, b.Update(l, 24))
|
assert.True(t, b.Update(l, 14))
|
||||||
assert.EqualValues(t, 25, b.current)
|
assert.EqualValues(t, 15, b.current)
|
||||||
g = []bool{false, false, false, false, false, false, false, false, true, true, false, false, false, false, false, false}
|
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
||||||
assert.Equal(t, g, b.snapshot())
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Mark 5, not allowed because 5 <= current-length (25-16=9).
|
// Mark 5, which is not allowed because it is not in the window
|
||||||
assert.False(t, b.Check(l, 5))
|
assert.False(t, b.Check(l, 5))
|
||||||
assert.False(t, b.Update(l, 5))
|
assert.False(t, b.Update(l, 5))
|
||||||
assert.EqualValues(t, 25, b.current)
|
assert.EqualValues(t, 15, b.current)
|
||||||
g = []bool{false, false, false, false, false, false, false, false, true, true, false, false, false, false, false, false}
|
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
||||||
assert.Equal(t, g, b.snapshot())
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Make sure we handle wrapping around once to the same slot. With
|
// make sure we handle wrapping around once to the current position
|
||||||
// length=16, packets 1 and 17 share slot 1.
|
b = NewBits(10)
|
||||||
b = NewBits(16)
|
|
||||||
assert.True(t, b.Update(l, 1))
|
assert.True(t, b.Update(l, 1))
|
||||||
assert.True(t, b.Update(l, 17))
|
assert.True(t, b.Update(l, 11))
|
||||||
assert.Equal(t, []bool{false, true, false, false, false, false, false, false, false, false, false, false, false, false, false, false}, b.snapshot())
|
assert.Equal(t, []bool{false, true, false, false, false, false, false, false, false, false}, b.bits)
|
||||||
|
|
||||||
// Walk through a few windows in order
|
// Walk through a few windows in order
|
||||||
b = NewBits(16)
|
b = NewBits(10)
|
||||||
for i := uint64(1); i <= 100; i++ {
|
for i := uint64(1); i <= 100; i++ {
|
||||||
assert.True(t, b.Check(l, i), "Error while checking %v", i)
|
assert.True(t, b.Check(l, i), "Error while checking %v", i)
|
||||||
assert.True(t, b.Update(l, i), "Error while updating %v", i)
|
assert.True(t, b.Update(l, i), "Error while updating %v", i)
|
||||||
@@ -90,31 +72,24 @@ func TestBits(t *testing.T) {
|
|||||||
|
|
||||||
func TestBitsLargeJumps(t *testing.T) {
|
func TestBitsLargeJumps(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
|
b := NewBits(10)
|
||||||
// length=16. Update(55) from current=0:
|
|
||||||
// warmup, per-bit loop sees no n>16 with unset bits (slot 0 was set by
|
|
||||||
// NewBits and gets re-evaluated when n=16; n=16 is not strictly > 16),
|
|
||||||
// so the loop contributes 0. The jump exceeds the window so we record
|
|
||||||
// 55 - 0 - 16 = 39 packets fell out the back.
|
|
||||||
b := NewBits(16)
|
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
assert.True(t, b.Update(l, 55))
|
|
||||||
assert.Equal(t, int64(39), b.lostCounter.Count())
|
|
||||||
|
|
||||||
// Update(100): clears 16 slots starting at slot 56%16=8. Only slot 7 (for
|
b = NewBits(10)
|
||||||
// packet 55) was set, so 16 - 1 = 15 evicted slots had unset bits.
|
b.lostCounter.Clear()
|
||||||
// Plus 100 - 55 - 16 = 29 packets fell past the window. Total 44.
|
assert.True(t, b.Update(l, 55)) // We saw packet 55 and can still track 45,46,47,48,49,50,51,52,53,54
|
||||||
assert.True(t, b.Update(l, 100))
|
assert.Equal(t, int64(45), b.lostCounter.Count())
|
||||||
assert.Equal(t, int64(39+44), b.lostCounter.Count())
|
|
||||||
|
|
||||||
// Update(200): same shape: 16 - 1 = 15 evicted unset, plus 200 - 100 - 16 = 84 past window. Total 99.
|
assert.True(t, b.Update(l, 100)) // We saw packet 55 and 100 and can still track 90,91,92,93,94,95,96,97,98,99
|
||||||
assert.True(t, b.Update(l, 200))
|
assert.Equal(t, int64(89), b.lostCounter.Count())
|
||||||
assert.Equal(t, int64(39+44+99), b.lostCounter.Count())
|
|
||||||
|
assert.True(t, b.Update(l, 200)) // We saw packet 55, 100, and 200 and can still track 190,191,192,193,194,195,196,197,198,199
|
||||||
|
assert.Equal(t, int64(188), b.lostCounter.Count())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBitsDupeCounter(t *testing.T) {
|
func TestBitsDupeCounter(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
b := NewBits(16)
|
b := NewBits(10)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
@@ -139,117 +114,120 @@ func TestBitsDupeCounter(t *testing.T) {
|
|||||||
|
|
||||||
func TestBitsOutOfWindowCounter(t *testing.T) {
|
func TestBitsOutOfWindowCounter(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
b := NewBits(16)
|
b := NewBits(10)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
|
|
||||||
// Jump to 20 (warmup branch + 4 past-window packets).
|
|
||||||
assert.True(t, b.Update(l, 20))
|
assert.True(t, b.Update(l, 20))
|
||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||||
|
|
||||||
// 9 single-step advances, each evicts a slot whose bit was cleared during
|
assert.True(t, b.Update(l, 21))
|
||||||
// the jump above and whose value was never seen, so each contributes 1
|
assert.True(t, b.Update(l, 22))
|
||||||
// to lostCounter.
|
assert.True(t, b.Update(l, 23))
|
||||||
for n := uint64(21); n <= 29; n++ {
|
assert.True(t, b.Update(l, 24))
|
||||||
assert.True(t, b.Update(l, n))
|
assert.True(t, b.Update(l, 25))
|
||||||
}
|
assert.True(t, b.Update(l, 26))
|
||||||
|
assert.True(t, b.Update(l, 27))
|
||||||
|
assert.True(t, b.Update(l, 28))
|
||||||
|
assert.True(t, b.Update(l, 29))
|
||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||||
|
|
||||||
// 0 is below current-length (29-16=13) so it falls outside the window.
|
|
||||||
assert.False(t, b.Update(l, 0))
|
assert.False(t, b.Update(l, 0))
|
||||||
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
||||||
|
|
||||||
// 4 from the Update(20) jump + 9 from 21..29.
|
assert.Equal(t, int64(19), b.lostCounter.Count()) // packet 0 wasn't lost
|
||||||
assert.Equal(t, int64(13), b.lostCounter.Count())
|
|
||||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||||
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBitsLostCounter(t *testing.T) {
|
func TestBitsLostCounter(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
b := NewBits(16)
|
b := NewBits(10)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
|
|
||||||
// Walk 20..29 like the original, just with a bigger window. Same
|
assert.True(t, b.Update(l, 20))
|
||||||
// reasoning as TestBitsOutOfWindowCounter: 4 past-window from Update(20),
|
assert.True(t, b.Update(l, 21))
|
||||||
// then 9 more from the unit advances.
|
assert.True(t, b.Update(l, 22))
|
||||||
for n := uint64(20); n <= 29; n++ {
|
assert.True(t, b.Update(l, 23))
|
||||||
assert.True(t, b.Update(l, n))
|
assert.True(t, b.Update(l, 24))
|
||||||
}
|
assert.True(t, b.Update(l, 25))
|
||||||
assert.Equal(t, int64(13), b.lostCounter.Count())
|
assert.True(t, b.Update(l, 26))
|
||||||
|
assert.True(t, b.Update(l, 27))
|
||||||
|
assert.True(t, b.Update(l, 28))
|
||||||
|
assert.True(t, b.Update(l, 29))
|
||||||
|
assert.Equal(t, int64(19), b.lostCounter.Count()) // packet 0 wasn't lost
|
||||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||||
|
|
||||||
b = NewBits(16)
|
b = NewBits(10)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
|
|
||||||
// Update(15) clears the warmup window (no lost), sets slot 15.
|
assert.True(t, b.Update(l, 9))
|
||||||
assert.True(t, b.Update(l, 15))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
|
// 10 will set 0 index, 0 was already set, no lost packets
|
||||||
// Update(16): slot 0 was already set (NewBits seeded it), and 16 is not
|
assert.True(t, b.Update(l, 10))
|
||||||
// strictly > length, so nothing is recorded as lost.
|
|
||||||
assert.True(t, b.Update(l, 16))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
|
// 11 will set 1 index, 1 was missed, we should see 1 packet lost
|
||||||
// Update(17): we jumped straight from 0 to 15, so slot 1 was cleared
|
assert.True(t, b.Update(l, 11))
|
||||||
// (and never re-set). 17 > 16 is past warmup, so packet 1 is recorded lost.
|
|
||||||
assert.True(t, b.Update(l, 17))
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
|
// Now let's fill in the window, should end up with 8 lost packets
|
||||||
|
assert.True(t, b.Update(l, 12))
|
||||||
|
assert.True(t, b.Update(l, 13))
|
||||||
|
assert.True(t, b.Update(l, 14))
|
||||||
|
assert.True(t, b.Update(l, 15))
|
||||||
|
assert.True(t, b.Update(l, 16))
|
||||||
|
assert.True(t, b.Update(l, 17))
|
||||||
|
assert.True(t, b.Update(l, 18))
|
||||||
|
assert.True(t, b.Update(l, 19))
|
||||||
|
assert.Equal(t, int64(8), b.lostCounter.Count())
|
||||||
|
|
||||||
// Fill in 18..30 in single steps. Each i evicts slot i%16. Slots 2..14
|
// Jump ahead by a window size
|
||||||
// were all cleared during Update(15), and we never re-set any of them,
|
assert.True(t, b.Update(l, 29))
|
||||||
// so each i in 18..30 is a fresh lost packet — 13 more.
|
assert.Equal(t, int64(8), b.lostCounter.Count())
|
||||||
for n := uint64(18); n <= 30; n++ {
|
// Now lets walk ahead normally through the window, the missed packets should fill in
|
||||||
assert.True(t, b.Update(l, n))
|
assert.True(t, b.Update(l, 30))
|
||||||
}
|
assert.True(t, b.Update(l, 31))
|
||||||
assert.Equal(t, int64(14), b.lostCounter.Count())
|
assert.True(t, b.Update(l, 32))
|
||||||
|
assert.True(t, b.Update(l, 33))
|
||||||
|
assert.True(t, b.Update(l, 34))
|
||||||
|
assert.True(t, b.Update(l, 35))
|
||||||
|
assert.True(t, b.Update(l, 36))
|
||||||
|
assert.True(t, b.Update(l, 37))
|
||||||
|
assert.True(t, b.Update(l, 38))
|
||||||
|
// 39 packets tracked, 22 seen, 17 lost
|
||||||
|
assert.Equal(t, int64(17), b.lostCounter.Count())
|
||||||
|
|
||||||
// Jump ahead by exactly one window size.
|
// Jump ahead by 2 windows, should have recording 1 full window missing
|
||||||
assert.True(t, b.Update(l, 46))
|
assert.True(t, b.Update(l, 58))
|
||||||
// end = min(46, 30+16) = 46, count = 16, all slots cleared. Before the
|
assert.Equal(t, int64(27), b.lostCounter.Count())
|
||||||
// jump every slot 0..15 had been set (Update(15), (16), (17), 18..30),
|
// Now lets walk ahead normally through the window, the missed packets should fill in from this window
|
||||||
// so wasSet=16 and 46 == current+length means no past-window slack:
|
assert.True(t, b.Update(l, 59))
|
||||||
// lost contribution = 0.
|
assert.True(t, b.Update(l, 60))
|
||||||
assert.Equal(t, int64(14), b.lostCounter.Count())
|
assert.True(t, b.Update(l, 61))
|
||||||
|
assert.True(t, b.Update(l, 62))
|
||||||
// Walk 47..55. The Update(46) jump cleared every slot, so only slot 14
|
assert.True(t, b.Update(l, 63))
|
||||||
// (for packet 46) is set when we start. Each subsequent unit step lands
|
assert.True(t, b.Update(l, 64))
|
||||||
// on a slot that was cleared and is past warmup, so it counts as lost.
|
assert.True(t, b.Update(l, 65))
|
||||||
// 9 more = 23.
|
assert.True(t, b.Update(l, 66))
|
||||||
for n := uint64(47); n <= 55; n++ {
|
assert.True(t, b.Update(l, 67))
|
||||||
assert.True(t, b.Update(l, n))
|
// 68 packets tracked, 32 seen, 36 missed
|
||||||
}
|
assert.Equal(t, int64(36), b.lostCounter.Count())
|
||||||
assert.Equal(t, int64(23), b.lostCounter.Count())
|
|
||||||
|
|
||||||
// Jump ahead by two windows: clears the window plus past-window loss.
|
|
||||||
assert.True(t, b.Update(l, 87))
|
|
||||||
// current=55, length=16. end = min(87, 71) = 71. count=16, all slots
|
|
||||||
// cleared. Slots set before the clear are slots 14,15,0..7 (10 total).
|
|
||||||
// Lost from clear = 16 - 10 = 6. Past window: 87 - 55 - 16 = 16. +22.
|
|
||||||
assert.Equal(t, int64(45), b.lostCounter.Count())
|
|
||||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBitsLostCounterIssue1(t *testing.T) {
|
func TestBitsLostCounterIssue1(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
b := NewBits(16)
|
b := NewBits(10)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
|
|
||||||
// Receive 4, backfill 1, then 9, 2, 3, 5, 6, 7 (skip 8), 10, 11, 14.
|
|
||||||
// Then jump to 25 — slot 25%16=9 is being evicted, but it had been set
|
|
||||||
// (we received packet 9), so no spurious lost increment. The original
|
|
||||||
// regression was about double-counting a missing packet when its slot
|
|
||||||
// got cleared on a jump. With the jump path now using clearRange's
|
|
||||||
// word-level wasSet count, the same semantics hold.
|
|
||||||
assert.True(t, b.Update(l, 4))
|
assert.True(t, b.Update(l, 4))
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 1))
|
assert.True(t, b.Update(l, 1))
|
||||||
@@ -266,7 +244,7 @@ func TestBitsLostCounterIssue1(t *testing.T) {
|
|||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 7))
|
assert.True(t, b.Update(l, 7))
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
// Skip packet 8.
|
// assert.True(t, b.Update(l, 8))
|
||||||
assert.True(t, b.Update(l, 10))
|
assert.True(t, b.Update(l, 10))
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 11))
|
assert.True(t, b.Update(l, 11))
|
||||||
@@ -274,23 +252,9 @@ func TestBitsLostCounterIssue1(t *testing.T) {
|
|||||||
|
|
||||||
assert.True(t, b.Update(l, 14))
|
assert.True(t, b.Update(l, 14))
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
|
// Issue seems to be here, we reset missing packet 8 to false here and don't increment the lost counter
|
||||||
// Jump to 25. With length=16, slot 25%16=9 corresponds to packet 9
|
assert.True(t, b.Update(l, 19))
|
||||||
// (which we DID receive), so its bit is set and no lost++ from that
|
|
||||||
// eviction. The trace below shows the only loss is packet 8.
|
|
||||||
assert.True(t, b.Update(l, 25))
|
|
||||||
// current was 14, i=25. end=min(25,30)=25. count=11. startPos=15.
|
|
||||||
// steady? current=14<16, so warmup branch: per-bit n=15..25, count those
|
|
||||||
// with !get(n) AND n>16. n=17..25 are >16. Among slots 17%16=1..25%16=9
|
|
||||||
// did we set slots 1..9 (packets 1..9)? Yes for all but slot 8 (packet 8
|
|
||||||
// was skipped). n=24 maps to slot 8 which is FALSE → lost++. All other
|
|
||||||
// n in 17..25 map to slots that are set. n=16 is not strictly > 16. So
|
|
||||||
// lost = 1.
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
|
|
||||||
// Fill in 12, 13, 15, 16. Each is below current=25 (in-window). 16 must
|
|
||||||
// recheck slot 0 — it was set by NewBits and then cleared by the
|
|
||||||
// Update(25) jump, so 16 backfills cleanly.
|
|
||||||
assert.True(t, b.Update(l, 12))
|
assert.True(t, b.Update(l, 12))
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 13))
|
assert.True(t, b.Update(l, 13))
|
||||||
@@ -299,140 +263,29 @@ func TestBitsLostCounterIssue1(t *testing.T) {
|
|||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 16))
|
assert.True(t, b.Update(l, 16))
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
|
assert.True(t, b.Update(l, 17))
|
||||||
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
|
assert.True(t, b.Update(l, 18))
|
||||||
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
|
assert.True(t, b.Update(l, 20))
|
||||||
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
|
assert.True(t, b.Update(l, 21))
|
||||||
|
|
||||||
// We missed packet 8 above and that loss is still recorded once, never
|
// We missed packet 8 above
|
||||||
// double-counted, never zeroed.
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestBitsWarmupOvershoot exercises the jump path's warmup arm with an
|
func BenchmarkBits(b *testing.B) {
|
||||||
// overshoot past one full window. NewBits leaves current=0 with only slot 0
|
z := NewBits(10)
|
||||||
// "set" by the marker. Jumping straight to length+k must (a) clear every
|
|
||||||
// slot the jump straddles, (b) count only past-window slack (not the
|
|
||||||
// in-window slots, which never had a "lost" tenant during warmup), and
|
|
||||||
// (c) leave the cursor at the new counter so subsequent unit advances
|
|
||||||
// count from steady state. The marker bit at slot 0 is irrelevant once
|
|
||||||
// current >= length.
|
|
||||||
func TestBitsWarmupOvershoot(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
b := NewBits(16)
|
|
||||||
b.lostCounter.Clear()
|
|
||||||
|
|
||||||
// Jump from current=0 to i=20 (length=16, overshoot=4).
|
|
||||||
// Warmup arm: counts slots in [1..16] where bit unset and n>length.
|
|
||||||
// Only n=16 was unset and >length: but slot 16%16=0 is the marker,
|
|
||||||
// so b.get(16) reads bits[0]=1 and skips. Result: 0 lost from the loop.
|
|
||||||
// Past-window: i - current - length = 20 - 0 - 16 = 4 lost.
|
|
||||||
assert.True(t, b.Update(l, 20))
|
|
||||||
assert.Equal(t, int64(4), b.lostCounter.Count())
|
|
||||||
assert.Equal(t, uint64(20), b.current)
|
|
||||||
|
|
||||||
// Steady state now (current=20 >= length=16). Unit advance to 21
|
|
||||||
// stomps slot 21%16=5, which was cleared by the jump and not reset,
|
|
||||||
// so this is +1 lost.
|
|
||||||
assert.True(t, b.Update(l, 21))
|
|
||||||
assert.Equal(t, int64(5), b.lostCounter.Count())
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestBitsCheckAcrossWarmupBoundary pins the underflow trick in Check's
|
|
||||||
// in-window clause. While in warmup, b.current-b.length underflows uint64
|
|
||||||
// to a huge value so the first OR-clause is always false; the second
|
|
||||||
// clause (i < length && current < length) carries the in-window check.
|
|
||||||
// Once current >= length the regimes flip cleanly.
|
|
||||||
func TestBitsCheckAcrossWarmupBoundary(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
b := NewBits(16)
|
|
||||||
|
|
||||||
// Warmup: current=0. Check(0) must read the marker (set) and return false.
|
|
||||||
assert.False(t, b.Check(l, 0), "marker slot should look already-received")
|
|
||||||
// Warmup: any 0 < i < length is in-window and unset → accepted.
|
|
||||||
for i := uint64(1); i < 16; i++ {
|
|
||||||
assert.True(t, b.Check(l, i), "warmup in-window i=%d should be accepted", i)
|
|
||||||
}
|
|
||||||
// Warmup: i >= length but > current is "next number" so accepted.
|
|
||||||
assert.True(t, b.Check(l, 16))
|
|
||||||
assert.True(t, b.Check(l, 1_000_000))
|
|
||||||
|
|
||||||
// Cross into steady state.
|
|
||||||
assert.True(t, b.Update(l, 100))
|
|
||||||
// Now current=100, length=16. In-window range is [85..100].
|
|
||||||
// 84 is just outside: the underflow clause activates; 84 > 100-16=84 is false.
|
|
||||||
// And the warmup clause is false (current >= length). So out of window.
|
|
||||||
assert.False(t, b.Check(l, 84))
|
|
||||||
// 85 sits at the boundary. 85 > 84 is true → in window, unset → accept.
|
|
||||||
assert.True(t, b.Check(l, 85))
|
|
||||||
// 100 is current itself; not strictly greater, in-window, but already set.
|
|
||||||
assert.False(t, b.Check(l, 100))
|
|
||||||
// Way out: clearly out of window.
|
|
||||||
assert.False(t, b.Check(l, 50))
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestBitsMarkerInvariant verifies the seeded bits[0]=1 marker behaves
|
|
||||||
// correctly across warmup and beyond. Update should never clear the marker
|
|
||||||
// during warmup (clearRange skips position 0 when startPos=1), and once
|
|
||||||
// current >= length the marker is no longer consulted by Check/Update on
|
|
||||||
// the live path — but it must still report counter 0 as a duplicate while
|
|
||||||
// we are in warmup.
|
|
||||||
func TestBitsMarkerInvariant(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
b := NewBits(8)
|
|
||||||
|
|
||||||
// Counter 0 is the seeded marker; Check sees it as already received.
|
|
||||||
assert.False(t, b.Check(l, 0))
|
|
||||||
// Update(0) at current=0 hits the duplicate branch.
|
|
||||||
b.dupeCounter.Clear()
|
|
||||||
assert.False(t, b.Update(l, 0))
|
|
||||||
assert.Equal(t, int64(1), b.dupeCounter.Count())
|
|
||||||
|
|
||||||
// Walk forward through warmup; the marker must remain set.
|
|
||||||
for n := uint64(1); n <= 7; n++ {
|
|
||||||
assert.True(t, b.Update(l, n))
|
|
||||||
}
|
|
||||||
// Position 0 (the marker) should still read as set because we never
|
|
||||||
// cleared it; Update(0) still looks like a duplicate.
|
|
||||||
assert.False(t, b.Check(l, 0))
|
|
||||||
|
|
||||||
// Cross into steady state with a unit advance to 8: pos=0, evicts the
|
|
||||||
// marker bit. The lost-counter guard (i > b.length) is false (8 == 8),
|
|
||||||
// so this advance does NOT charge a lost packet — exactly what the
|
|
||||||
// marker is there to prevent.
|
|
||||||
b.lostCounter.Clear()
|
|
||||||
assert.True(t, b.Update(l, 8))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
// The slot at pos 0 is now occupied by counter 8.
|
|
||||||
assert.False(t, b.Check(l, 8))
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkBitsUpdateInOrder is the steady-state hot path: each call is
|
|
||||||
// i == current+1.
|
|
||||||
func BenchmarkBitsUpdateInOrder(b *testing.B) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
z := NewBits(16384)
|
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
z.Update(l, uint64(n)+1)
|
for i := range z.bits {
|
||||||
}
|
z.bits[i] = true
|
||||||
}
|
}
|
||||||
|
for i := range z.bits {
|
||||||
|
z.bits[i] = false
|
||||||
|
}
|
||||||
|
|
||||||
// BenchmarkBitsUpdateReorder simulates light reorder within the window:
|
|
||||||
// every other packet arrives one slot behind its predecessor (forces the
|
|
||||||
// in-window backfill branch).
|
|
||||||
func BenchmarkBitsUpdateReorder(b *testing.B) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
z := NewBits(16384)
|
|
||||||
for n := 0; n < b.N; n++ {
|
|
||||||
base := uint64(n) * 2
|
|
||||||
z.Update(l, base+2)
|
|
||||||
z.Update(l, base+1)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkBitsUpdateLargeJumps stresses the clearRange word-level path.
|
|
||||||
func BenchmarkBitsUpdateLargeJumps(b *testing.B) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
z := NewBits(16384)
|
|
||||||
for n := 0; n < b.N; n++ {
|
|
||||||
z.Update(l, uint64(n+1)*1000)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -217,10 +217,6 @@ func (ncp *CAPool) verify(c Certificate, now time.Time, certFp string, signerFp
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
if signer.Certificate.Curve() != c.Curve() {
|
|
||||||
return nil, ErrCurveMismatch
|
|
||||||
}
|
|
||||||
|
|
||||||
if signer.Certificate.Expired(now) {
|
if signer.Certificate.Expired(now) {
|
||||||
return nil, ErrRootExpired
|
return nil, ErrRootExpired
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -654,31 +654,3 @@ func TestCertificateV2_Verify_Subnets(t *testing.T) {
|
|||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
_, err = caPool.VerifyCertificate(time.Now(), c)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCertificateV2_CurveMismatch(t *testing.T) {
|
|
||||||
caIp1 := mustParsePrefixUnmapped("10.0.0.0/16")
|
|
||||||
caIp2 := mustParsePrefixUnmapped("192.168.0.0/24")
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version2, Curve_P256, time.Now(), time.Now().Add(10*time.Minute), []netip.Prefix{caIp1, caIp2}, nil, []string{"test"})
|
|
||||||
|
|
||||||
caPem, err := ca.MarshalPEM()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
caPool := NewCAPool()
|
|
||||||
b, err := caPool.AddCAFromPEM(caPem)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
|
|
||||||
// ip is outside the network
|
|
||||||
cIp1 := mustParsePrefixUnmapped("10.0.0.1/24")
|
|
||||||
c, _, _, _ := NewTestCert(Version2, Curve_P256, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1}, nil, []string{"test"})
|
|
||||||
|
|
||||||
fp, _ := c.Fingerprint()
|
|
||||||
_, err = caPool.verify(c, time.Now(), fp, c.Issuer())
|
|
||||||
require.NoError(t, err)
|
|
||||||
//
|
|
||||||
c2 := c.(*certificateV2)
|
|
||||||
c2.curve = Curve_CURVE25519
|
|
||||||
fp, _ = c.Fingerprint()
|
|
||||||
_, err = caPool.verify(c, time.Now(), fp, c.Issuer())
|
|
||||||
require.Error(t, err)
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -112,9 +112,6 @@ func (c *certificateV1) CheckSignature(key []byte) bool {
|
|||||||
}
|
}
|
||||||
switch c.details.curve {
|
switch c.details.curve {
|
||||||
case Curve_CURVE25519:
|
case Curve_CURVE25519:
|
||||||
if len(key) != ed25519.PublicKeySize {
|
|
||||||
return false //avoids a panic internal to ed25519
|
|
||||||
}
|
|
||||||
return ed25519.Verify(key, b, c.signature)
|
return ed25519.Verify(key, b, c.signature)
|
||||||
case Curve_P256:
|
case Curve_P256:
|
||||||
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
||||||
|
|||||||
@@ -151,9 +151,6 @@ func (c *certificateV2) CheckSignature(key []byte) bool {
|
|||||||
|
|
||||||
switch c.curve {
|
switch c.curve {
|
||||||
case Curve_CURVE25519:
|
case Curve_CURVE25519:
|
||||||
if len(key) != ed25519.PublicKeySize {
|
|
||||||
return false //avoids a panic internal to ed25519
|
|
||||||
}
|
|
||||||
return ed25519.Verify(key, b, c.signature)
|
return ed25519.Verify(key, b, c.signature)
|
||||||
case Curve_P256:
|
case Curve_P256:
|
||||||
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
||||||
|
|||||||
@@ -22,7 +22,6 @@ var (
|
|||||||
ErrCaNotFound = errors.New("could not find ca for the certificate")
|
ErrCaNotFound = errors.New("could not find ca for the certificate")
|
||||||
ErrUnknownVersion = errors.New("certificate version unrecognized")
|
ErrUnknownVersion = errors.New("certificate version unrecognized")
|
||||||
ErrCertPubkeyPresent = errors.New("certificate has unexpected pubkey present")
|
ErrCertPubkeyPresent = errors.New("certificate has unexpected pubkey present")
|
||||||
ErrCurveMismatch = errors.New("certificate curve does not match CA")
|
|
||||||
|
|
||||||
ErrInvalidPEMBlock = errors.New("input did not contain a valid PEM encoded block")
|
ErrInvalidPEMBlock = errors.New("input did not contain a valid PEM encoded block")
|
||||||
ErrInvalidPEMCertificateBanner = errors.New("bytes did not contain a proper certificate banner")
|
ErrInvalidPEMCertificateBanner = errors.New("bytes did not contain a proper certificate banner")
|
||||||
|
|||||||
+4
-10
@@ -13,12 +13,6 @@ 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
|
||||||
@@ -40,10 +34,10 @@ func NewTestCaCert(version Version, curve Curve, before, after time.Time, networ
|
|||||||
}
|
}
|
||||||
|
|
||||||
if before.IsZero() {
|
if before.IsZero() {
|
||||||
before = testCertNow.Add(time.Second * -60)
|
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
||||||
}
|
}
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = testCertNow.Add(time.Second * 60)
|
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
t := &TBSCertificate{
|
t := &TBSCertificate{
|
||||||
@@ -76,11 +70,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 = testCertNow.Add(time.Second * -60)
|
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = testCertNow.Add(time.Second * 60)
|
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(networks) == 0 {
|
if len(networks) == 0 {
|
||||||
|
|||||||
+2
-32
@@ -148,9 +148,6 @@ 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 {
|
||||||
@@ -159,10 +156,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:
|
case X25519PublicKeyBanner, Ed25519PublicKeyBanner:
|
||||||
expectedLen = 32
|
expectedLen = 32
|
||||||
curve = Curve_CURVE25519
|
curve = Curve_CURVE25519
|
||||||
case P256PublicKeyBanner:
|
case P256PublicKeyBanner, ECDSAP256PublicKeyBanner:
|
||||||
// Uncompressed
|
// Uncompressed
|
||||||
expectedLen = 65
|
expectedLen = 65
|
||||||
curve = Curve_P256
|
curve = Curve_P256
|
||||||
@@ -175,33 +172,6 @@ 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:
|
||||||
|
|||||||
+67
-87
@@ -255,6 +255,60 @@ 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-----
|
||||||
@@ -265,7 +319,7 @@ AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
|||||||
AAAAAAAAAAAAAAAAAAAAAAA=
|
AAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
-----END NEBULA P256 PUBLIC KEY-----
|
-----END NEBULA P256 PUBLIC KEY-----
|
||||||
`)
|
`)
|
||||||
signingKey := []byte(`# A signing key has the wrong scope for this function
|
oldPubP256Key := []byte(`# A good key
|
||||||
-----BEGIN NEBULA ECDSA P256 PUBLIC KEY-----
|
-----BEGIN NEBULA ECDSA P256 PUBLIC KEY-----
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
||||||
AAAAAAAAAAAAAAAAAAAAAAA=
|
AAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
@@ -286,118 +340,44 @@ AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
||||||
-END NEBULA X25519 PUBLIC KEY-----`)
|
-END NEBULA X25519 PUBLIC KEY-----`)
|
||||||
|
|
||||||
keyBundle := appendByteSlices(pubKey, pubP256Key, signingKey, shortKey, invalidBanner, invalidPem)
|
keyBundle := appendByteSlices(pubKey, pubP256Key, oldPubP256Key, shortKey, invalidBanner, invalidPem)
|
||||||
|
|
||||||
// X25519 key
|
// Success test case
|
||||||
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, signingKey, shortKey, invalidBanner, invalidPem))
|
assert.Equal(t, rest, appendByteSlices(pubP256Key, oldPubP256Key, shortKey, invalidBanner, invalidPem))
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
assert.Equal(t, Curve_CURVE25519, curve)
|
||||||
|
|
||||||
// P256 key
|
// Success test case
|
||||||
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(signingKey, shortKey, invalidBanner, invalidPem))
|
assert.Equal(t, rest, appendByteSlices(oldPubP256Key, shortKey, invalidBanner, invalidPem))
|
||||||
assert.Equal(t, Curve_P256, curve)
|
assert.Equal(t, Curve_P256, curve)
|
||||||
|
|
||||||
// Reject a signing public key (Ed25519/ECDSA banner)
|
// Success test case
|
||||||
k, rest, _, err = UnmarshalPublicKeyFromPEM(rest)
|
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
|
|
||||||
require.EqualError(t, err, "bytes did not contain a proper public key banner")
|
|
||||||
|
|
||||||
// Fail due to short key
|
|
||||||
k, rest, _, err = UnmarshalPublicKeyFromPEM(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 = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
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, _, 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)
|
assert.Len(t, k, 65)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, rest, appendByteSlices(ecdhKey, shortKey, invalidBanner, invalidPem))
|
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
|
||||||
assert.Equal(t, Curve_P256, curve)
|
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
|
// Fail due to short key
|
||||||
k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest)
|
k, rest, curve, 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, _, err = UnmarshalSigningPublicKeyFromPEM(rest)
|
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
||||||
assert.Nil(t, k)
|
assert.Nil(t, k)
|
||||||
require.EqualError(t, err, "bytes did not contain a proper Ed25519/ECDSA 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, _, err = UnmarshalSigningPublicKeyFromPEM(rest)
|
k, rest, curve, err = UnmarshalPublicKeyFromPEM(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")
|
||||||
|
|||||||
+4
-62
@@ -14,12 +14,6 @@ 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
|
||||||
@@ -41,10 +35,10 @@ func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Ti
|
|||||||
}
|
}
|
||||||
|
|
||||||
if before.IsZero() {
|
if before.IsZero() {
|
||||||
before = testCertNow.Add(time.Second * -60)
|
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
||||||
}
|
}
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = testCertNow.Add(time.Second * 60)
|
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
t := &cert.TBSCertificate{
|
t := &cert.TBSCertificate{
|
||||||
@@ -77,11 +71,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 = testCertNow.Add(time.Second * -60)
|
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = testCertNow.Add(time.Second * 60)
|
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
var pub, priv []byte
|
var pub, priv []byte
|
||||||
@@ -169,55 +163,3 @@ func P256Keypair() ([]byte, []byte) {
|
|||||||
pubkey := privkey.PublicKey()
|
pubkey := privkey.PublicKey()
|
||||||
return pubkey.Bytes(), privkey.Bytes()
|
return pubkey.Bytes(), privkey.Bytes()
|
||||||
}
|
}
|
||||||
|
|
||||||
// DummyCert is a minimal cert.Certificate implementation for testing error paths.
|
|
||||||
type DummyCert struct {
|
|
||||||
Version_ cert.Version
|
|
||||||
Curve_ cert.Curve
|
|
||||||
Groups_ []string
|
|
||||||
IsCA_ bool
|
|
||||||
Issuer_ string
|
|
||||||
Name_ string
|
|
||||||
Networks_ []netip.Prefix
|
|
||||||
NotAfter_ time.Time
|
|
||||||
NotBefore_ time.Time
|
|
||||||
PublicKey_ []byte
|
|
||||||
Signature_ []byte
|
|
||||||
UnsafeNetworks_ []netip.Prefix
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *DummyCert) Version() cert.Version { return d.Version_ }
|
|
||||||
func (d *DummyCert) Curve() cert.Curve { return d.Curve_ }
|
|
||||||
func (d *DummyCert) Groups() []string { return d.Groups_ }
|
|
||||||
func (d *DummyCert) IsCA() bool { return d.IsCA_ }
|
|
||||||
func (d *DummyCert) Issuer() string { return d.Issuer_ }
|
|
||||||
func (d *DummyCert) Name() string { return d.Name_ }
|
|
||||||
func (d *DummyCert) Networks() []netip.Prefix { return d.Networks_ }
|
|
||||||
func (d *DummyCert) NotAfter() time.Time { return d.NotAfter_ }
|
|
||||||
func (d *DummyCert) NotBefore() time.Time { return d.NotBefore_ }
|
|
||||||
func (d *DummyCert) PublicKey() []byte { return d.PublicKey_ }
|
|
||||||
func (d *DummyCert) Signature() []byte { return d.Signature_ }
|
|
||||||
func (d *DummyCert) UnsafeNetworks() []netip.Prefix { return d.UnsafeNetworks_ }
|
|
||||||
func (d *DummyCert) Fingerprint() (string, error) { return "", nil }
|
|
||||||
func (d *DummyCert) CheckSignature(key []byte) bool { return false }
|
|
||||||
func (d *DummyCert) MarshalForHandshakes() ([]byte, error) { return nil, nil }
|
|
||||||
func (d *DummyCert) MarshalPEM() ([]byte, error) { return nil, nil }
|
|
||||||
func (d *DummyCert) MarshalJSON() ([]byte, error) { return nil, nil }
|
|
||||||
func (d *DummyCert) Marshal() ([]byte, error) { return nil, nil }
|
|
||||||
func (d *DummyCert) String() string { return "dummy" }
|
|
||||||
func (d *DummyCert) Copy() cert.Certificate { return d }
|
|
||||||
func (d *DummyCert) VerifyPrivateKey(c cert.Curve, k []byte) error { return nil }
|
|
||||||
func (d *DummyCert) Expired(time.Time) bool { return false }
|
|
||||||
func (d *DummyCert) MarshalPublicKeyPEM() []byte { return nil }
|
|
||||||
func (d *DummyCert) PublicKeyPEM() []byte { return nil }
|
|
||||||
|
|
||||||
// NewTestCAPool creates a CAPool from the given CA certificates, panicking on error.
|
|
||||||
func NewTestCAPool(cas ...cert.Certificate) *cert.CAPool {
|
|
||||||
pool := cert.NewCAPool()
|
|
||||||
for _, ca := range cas {
|
|
||||||
if err := pool.AddCA(ca); err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return pool
|
|
||||||
}
|
|
||||||
|
|||||||
+8
-41
@@ -3,7 +3,6 @@ package main
|
|||||||
import (
|
import (
|
||||||
"crypto/ecdsa"
|
"crypto/ecdsa"
|
||||||
"crypto/elliptic"
|
"crypto/elliptic"
|
||||||
"crypto/fips140"
|
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
@@ -44,13 +43,6 @@ type caFlags struct {
|
|||||||
subnets *string
|
subnets *string
|
||||||
}
|
}
|
||||||
|
|
||||||
func defaultCurve() string {
|
|
||||||
if fips140.Enforced() {
|
|
||||||
return "P256"
|
|
||||||
}
|
|
||||||
return "25519"
|
|
||||||
}
|
|
||||||
|
|
||||||
func newCaFlags() *caFlags {
|
func newCaFlags() *caFlags {
|
||||||
cf := caFlags{set: flag.NewFlagSet("ca", flag.ContinueOnError)}
|
cf := caFlags{set: flag.NewFlagSet("ca", flag.ContinueOnError)}
|
||||||
cf.set.Usage = func() {}
|
cf.set.Usage = func() {}
|
||||||
@@ -67,7 +59,7 @@ func newCaFlags() *caFlags {
|
|||||||
cf.argonParallelism = cf.set.Uint("argon-parallelism", 4, "Optional: Argon2 parallelism parameter used for encrypted private key passphrase")
|
cf.argonParallelism = cf.set.Uint("argon-parallelism", 4, "Optional: Argon2 parallelism parameter used for encrypted private key passphrase")
|
||||||
cf.argonIterations = cf.set.Uint("argon-iterations", 1, "Optional: Argon2 iterations parameter used for encrypted private key passphrase")
|
cf.argonIterations = cf.set.Uint("argon-iterations", 1, "Optional: Argon2 iterations parameter used for encrypted private key passphrase")
|
||||||
cf.encryption = cf.set.Bool("encrypt", false, "Optional: prompt for passphrase and write out-key in an encrypted format")
|
cf.encryption = cf.set.Bool("encrypt", false, "Optional: prompt for passphrase and write out-key in an encrypted format")
|
||||||
cf.curve = cf.set.String("curve", defaultCurve(), "EdDSA/ECDSA Curve (25519, P256)")
|
cf.curve = cf.set.String("curve", "25519", "EdDSA/ECDSA Curve (25519, P256)")
|
||||||
cf.p11url = p11Flag(cf.set)
|
cf.p11url = p11Flag(cf.set)
|
||||||
|
|
||||||
cf.ips = cf.set.String("ips", "", "Deprecated, see -networks")
|
cf.ips = cf.set.String("ips", "", "Deprecated, see -networks")
|
||||||
@@ -105,19 +97,6 @@ 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
|
||||||
@@ -192,21 +171,12 @@ 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++ {
|
||||||
errOut.Write([]byte("Enter passphrase: "))
|
out.Write([]byte("Enter passphrase: "))
|
||||||
passphrase, err = pr.ReadPassword()
|
passphrase, err = pr.ReadPassword()
|
||||||
|
|
||||||
if err == ErrNoTerminal {
|
if err == ErrNoTerminal {
|
||||||
@@ -291,16 +261,14 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
Curve: curve,
|
Curve: curve,
|
||||||
}
|
}
|
||||||
|
|
||||||
if !isP11 && !isStdio(*cf.outKeyPath) {
|
if !isP11 {
|
||||||
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
|
||||||
@@ -326,7 +294,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 = writeOutput(*cf.outKeyPath, b, 0600, out)
|
err = os.WriteFile(*cf.outKeyPath, b, 0600)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -337,7 +305,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 = writeOutput(*cf.outCertPath, b, 0600, out)
|
err = os.WriteFile(*cf.outCertPath, b, 0600)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -348,7 +316,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 = writeOutput(*cf.outQRPath, b, 0600, out)
|
err = os.WriteFile(*cf.outQRPath, b, 0600)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -364,7 +332,6 @@ 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,7 +27,6 @@ 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"+
|
||||||
@@ -85,7 +84,7 @@ func Test_ca(t *testing.T) {
|
|||||||
err: nil,
|
err: nil,
|
||||||
}
|
}
|
||||||
|
|
||||||
pwPromptEB := "Enter passphrase: "
|
pwPromptOb := "Enter passphrase: "
|
||||||
|
|
||||||
// required args
|
// required args
|
||||||
assertHelpError(t, ca(
|
assertHelpError(t, ca(
|
||||||
@@ -169,8 +168,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.Empty(t, ob.String())
|
assert.Equal(t, pwPromptOb, ob.String())
|
||||||
assert.Equal(t, pwPromptEB, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// test encrypted key with passphrase environment variable
|
// test encrypted key with passphrase environment variable
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
@@ -208,8 +207,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.Empty(t, ob.String())
|
assert.Equal(t, pwPromptOb, ob.String())
|
||||||
assert.Equal(t, pwPromptEB, eb.String())
|
assert.Empty(t, 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())
|
||||||
@@ -218,8 +217,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.Empty(t, ob.String())
|
assert.Equal(t, strings.Repeat(pwPromptOb, 5), ob.String()) // prompts 5 times before giving up
|
||||||
assert.Equal(t, strings.Repeat(pwPromptEB, 5), eb.String()) // prompts 5 times before giving up
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// create valid cert/key for overwrite tests
|
// create valid cert/key for overwrite tests
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
@@ -248,67 +247,3 @@ 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
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,5 +0,0 @@
|
|||||||
//go:build fips140enforce
|
|
||||||
|
|
||||||
//go:debug fips140=only
|
|
||||||
|
|
||||||
package main
|
|
||||||
@@ -24,7 +24,7 @@ func newKeygenFlags() *keygenFlags {
|
|||||||
cf.set.Usage = func() {}
|
cf.set.Usage = func() {}
|
||||||
cf.outPubPath = cf.set.String("out-pub", "", "Required: path to write the public key to")
|
cf.outPubPath = cf.set.String("out-pub", "", "Required: path to write the public key to")
|
||||||
cf.outKeyPath = cf.set.String("out-key", "", "Required: path to write the private key to")
|
cf.outKeyPath = cf.set.String("out-key", "", "Required: path to write the private key to")
|
||||||
cf.curve = cf.set.String("curve", defaultCurve(), "ECDH Curve (25519, P256)")
|
cf.curve = cf.set.String("curve", "25519", "ECDH Curve (25519, P256)")
|
||||||
cf.p11url = p11Flag(cf.set)
|
cf.p11url = p11Flag(cf.set)
|
||||||
return &cf
|
return &cf
|
||||||
}
|
}
|
||||||
@@ -42,8 +42,6 @@ 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
|
||||||
@@ -71,14 +69,6 @@ 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 {
|
||||||
@@ -92,12 +82,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 = writeOutput(*cf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600, out)
|
err = os.WriteFile(*cf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600)
|
||||||
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 = writeOutput(*cf.outPubPath, cert.MarshalPublicKeyToPEM(curve, pub), 0600, out)
|
err = os.WriteFile(*cf.outPubPath, cert.MarshalPublicKeyToPEM(curve, pub), 0600)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -112,7 +102,6 @@ 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,7 +20,6 @@ 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"+
|
||||||
@@ -94,43 +93,3 @@ 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,9 +22,7 @@ func (pr StdinPasswordReader) ReadPassword() ([]byte, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
password, err := term.ReadPassword(int(os.Stdin.Fd()))
|
password, err := term.ReadPassword(int(os.Stdin.Fd()))
|
||||||
// Terminal echo is off while reading, so the user's Enter key does not
|
fmt.Println()
|
||||||
// produce a visible newline. Emit one on stderr to match the prompt.
|
|
||||||
fmt.Fprintln(os.Stderr)
|
|
||||||
|
|
||||||
return password, err
|
return password, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -40,23 +40,11 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
var claims ioClaims
|
rawCert, err := os.ReadFile(*pf.path)
|
||||||
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
|
||||||
@@ -69,13 +57,11 @@ 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 != "" {
|
||||||
@@ -93,7 +79,7 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
part++
|
part++
|
||||||
}
|
}
|
||||||
|
|
||||||
if *pf.json && !qrToStdout {
|
if *pf.json {
|
||||||
b, _ := json.Marshal(jsonCerts)
|
b, _ := json.Marshal(jsonCerts)
|
||||||
_, _ = out.Write(b)
|
_, _ = out.Write(b)
|
||||||
_, _ = out.Write([]byte("\n"))
|
_, _ = out.Write([]byte("\n"))
|
||||||
@@ -105,7 +91,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 = writeOutput(*pf.outQRPath, b, 0600, out)
|
err = os.WriteFile(*pf.outQRPath, b, 0600)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -121,7 +107,6 @@ 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,7 +25,6 @@ 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"+
|
||||||
@@ -179,44 +178,6 @@ 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
|
||||||
|
|||||||
+20
-42
@@ -85,9 +85,6 @@ 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
|
||||||
@@ -105,35 +102,13 @@ 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 = readInput("ca-key", *sf.caKeyPath, &claims)
|
rawCAKey, err := os.ReadFile(*sf.caKeyPath)
|
||||||
|
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -146,7 +121,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++ {
|
||||||
errOut.Write([]byte("Enter passphrase: "))
|
out.Write([]byte("Enter passphrase: "))
|
||||||
passphrase, err = pr.ReadPassword()
|
passphrase, err = pr.ReadPassword()
|
||||||
|
|
||||||
if errors.Is(err, ErrNoTerminal) {
|
if errors.Is(err, ErrNoTerminal) {
|
||||||
@@ -172,7 +147,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
rawCACert, err := readInput("ca-crt", *sf.caCertPath, &claims)
|
rawCACert, err := os.ReadFile(*sf.caCertPath)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -270,7 +245,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 := readInput("in-pub", *sf.inPubPath, &claims)
|
rawPub, err := os.ReadFile(*sf.inPubPath)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -291,10 +266,16 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
pub, rawPriv = newKeypair(curve)
|
pub, rawPriv = newKeypair(curve)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !isStdio(*sf.outCertPath) {
|
if *sf.outKeyPath == "" {
|
||||||
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
*sf.outKeyPath = *sf.name + ".key"
|
||||||
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
}
|
||||||
}
|
|
||||||
|
if *sf.outCertPath == "" {
|
||||||
|
*sf.outCertPath = *sf.name + ".crt"
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
||||||
|
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
var crts []cert.Certificate
|
var crts []cert.Certificate
|
||||||
@@ -379,13 +360,11 @@ 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 = writeOutput(*sf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600, out)
|
err = os.WriteFile(*sf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -400,7 +379,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
b = append(b, sb...)
|
b = append(b, sb...)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = writeOutput(*sf.outCertPath, b, 0600, out)
|
err = os.WriteFile(*sf.outCertPath, b, 0600)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -411,7 +390,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 = writeOutput(*sf.outQRPath, b, 0600, out)
|
err = os.WriteFile(*sf.outQRPath, b, 0600)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -461,7 +440,6 @@ 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,7 +27,6 @@ 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"+
|
||||||
@@ -377,18 +376,15 @@ 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.Empty(t, ob.String())
|
assert.Equal(t, "Enter passphrase: ", ob.String())
|
||||||
assert.Equal(t, "Enter passphrase: ", eb.String())
|
assert.Empty(t, 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", "")
|
||||||
|
|
||||||
@@ -399,8 +395,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.Empty(t, ob.String())
|
assert.Equal(t, "Enter passphrase: ", ob.String())
|
||||||
assert.Equal(t, "Enter passphrase: ", eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// test with the wrong password in environment
|
// test with the wrong password in environment
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
@@ -420,8 +416,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.Empty(t, ob.String())
|
assert.Equal(t, "Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: ", ob.String())
|
||||||
assert.Equal(t, "Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: ", eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// test an error condition
|
// test an error condition
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
@@ -429,106 +425,6 @@ 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.Empty(t, ob.String())
|
assert.Equal(t, "Enter passphrase: ", ob.String())
|
||||||
assert.Equal(t, "Enter passphrase: ", eb.String())
|
assert.Empty(t, 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())
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,117 +0,0 @@
|
|||||||
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
|
|
||||||
}
|
|
||||||
@@ -1,167 +0,0 @@
|
|||||||
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,26 +39,18 @@ func verify(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
var claims ioClaims
|
caFile, err := os.Open(*vf.caPath)
|
||||||
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 caReader.Close()
|
defer caFile.Close()
|
||||||
|
|
||||||
caPool, err := cert.NewCAPoolFromPEMReader(caReader)
|
caPool, err := cert.NewCAPoolFromPEMReader(caFile)
|
||||||
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 := readInput("crt", *vf.certPath, &claims)
|
rawCert, err := os.ReadFile(*vf.certPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("unable to read crt: %w", err)
|
return fmt.Errorf("unable to read crt: %w", err)
|
||||||
}
|
}
|
||||||
@@ -93,7 +85,6 @@ 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,7 +23,6 @@ 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"+
|
||||||
@@ -123,46 +122,3 @@ 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`)
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,5 +0,0 @@
|
|||||||
//go:build fips140enforce
|
|
||||||
|
|
||||||
//go:debug fips140=only
|
|
||||||
|
|
||||||
package main
|
|
||||||
@@ -3,15 +3,8 @@
|
|||||||
|
|
||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import "github.com/sirupsen/logrus"
|
||||||
"log/slog"
|
|
||||||
"os"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/logging"
|
func HookLogger(l *logrus.Logger) {
|
||||||
)
|
// Do nothing, let the logs flow to stdout/stderr
|
||||||
|
|
||||||
// newPlatformLogger returns a *slog.Logger that writes to stdout. Non-Windows
|
|
||||||
// platforms have no special sink to integrate with.
|
|
||||||
func newPlatformLogger() *slog.Logger {
|
|
||||||
return logging.NewLogger(os.Stdout)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,86 +1,54 @@
|
|||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"fmt"
|
||||||
"log/slog"
|
"io/ioutil"
|
||||||
"strings"
|
"os"
|
||||||
"sync"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/logging"
|
"github.com/kardianos/service"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
)
|
)
|
||||||
|
|
||||||
// newPlatformLogger returns a *slog.Logger that routes every log record
|
// HookLogger routes the logrus logs through the service logger so that they end up in the Windows Event Viewer
|
||||||
// through the Windows service logger so records end up in the Windows
|
// logrus output will be discarded
|
||||||
// Event Log. All the heavy lifting (level management, format swap,
|
func HookLogger(l *logrus.Logger) {
|
||||||
// timestamp toggle, WithAttrs/WithGroup) comes from logging.NewHandler;
|
l.AddHook(newLogHook(logger))
|
||||||
// this file only contributes:
|
l.SetOutput(ioutil.Discard)
|
||||||
//
|
|
||||||
// - an io.Writer that forwards each formatted line to the service
|
|
||||||
// logger at the current record's Event Log severity, and
|
|
||||||
// - a thin severityTag that embeds *logging.Handler and overrides
|
|
||||||
// only Handle / WithAttrs / WithGroup, so Event Viewer's severity
|
|
||||||
// column and severity-based filters keep working the way they did
|
|
||||||
// before the slog migration.
|
|
||||||
//
|
|
||||||
// Format (text vs json) is carried by the embedded *logging.Handler, so
|
|
||||||
// logging.format: json in config still produces JSON lines in Event
|
|
||||||
// Viewer, same as the pre-slog logrus setup.
|
|
||||||
func newPlatformLogger() *slog.Logger {
|
|
||||||
w := &eventLogWriter{}
|
|
||||||
return slog.New(&severityTag{Handler: logging.NewHandler(w), w: w})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// eventLogWriter forwards slog-formatted lines to the Windows service
|
type logHook struct {
|
||||||
// logger at the severity most recently stashed by severityTag.Handle.
|
sl service.Logger
|
||||||
// The mutex serializes the stash + inner.Handle + Write cycle per record
|
|
||||||
// across all concurrent goroutines; slog's builtin text/json handlers
|
|
||||||
// each hold their own mutex around Write, but that only protects the
|
|
||||||
// Write call itself, not our stash-then-handle sequence.
|
|
||||||
type eventLogWriter struct {
|
|
||||||
mu sync.Mutex
|
|
||||||
level slog.Level
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *eventLogWriter) Write(p []byte) (int, error) {
|
func newLogHook(sl service.Logger) *logHook {
|
||||||
line := strings.TrimRight(string(p), "\n")
|
return &logHook{sl: sl}
|
||||||
switch {
|
}
|
||||||
case w.level >= slog.LevelError:
|
|
||||||
return len(p), logger.Error(line)
|
func (h *logHook) Fire(entry *logrus.Entry) error {
|
||||||
case w.level >= slog.LevelWarn:
|
line, err := entry.String()
|
||||||
return len(p), logger.Warning(line)
|
if err != nil {
|
||||||
|
fmt.Fprintf(os.Stderr, "Unable to read entry, %v", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
switch entry.Level {
|
||||||
|
case logrus.PanicLevel:
|
||||||
|
return h.sl.Error(line)
|
||||||
|
case logrus.FatalLevel:
|
||||||
|
return h.sl.Error(line)
|
||||||
|
case logrus.ErrorLevel:
|
||||||
|
return h.sl.Error(line)
|
||||||
|
case logrus.WarnLevel:
|
||||||
|
return h.sl.Warning(line)
|
||||||
|
case logrus.InfoLevel:
|
||||||
|
return h.sl.Info(line)
|
||||||
|
case logrus.DebugLevel:
|
||||||
|
return h.sl.Info(line)
|
||||||
default:
|
default:
|
||||||
return len(p), logger.Info(line)
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// severityTag embeds *logging.Handler to pick up everything it does for
|
func (h *logHook) Levels() []logrus.Level {
|
||||||
// free (Enabled, SetLevel, GetLevel, SetFormat, GetFormat,
|
return logrus.AllLevels
|
||||||
// SetDisableTimestamp) and overrides only Handle / WithAttrs / WithGroup
|
|
||||||
// so each record's slog.Level is stashed on the writer before formatting
|
|
||||||
// and so derived handlers stay wrapped as severityTag rather than
|
|
||||||
// downgrading to bare *logging.Handler.
|
|
||||||
type severityTag struct {
|
|
||||||
*logging.Handler
|
|
||||||
w *eventLogWriter
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *severityTag) Handle(ctx context.Context, r slog.Record) error {
|
|
||||||
s.w.mu.Lock()
|
|
||||||
defer s.w.mu.Unlock()
|
|
||||||
s.w.level = r.Level
|
|
||||||
return s.Handler.Handle(ctx, r)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *severityTag) WithAttrs(attrs []slog.Attr) slog.Handler {
|
|
||||||
if len(attrs) == 0 {
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
return &severityTag{Handler: s.Handler.WithAttrs(attrs).(*logging.Handler), w: s.w}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *severityTag) WithGroup(name string) slog.Handler {
|
|
||||||
if name == "" {
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
return &severityTag{Handler: s.Handler.WithGroup(name).(*logging.Handler), w: s.w}
|
|
||||||
}
|
}
|
||||||
|
|||||||
+12
-28
@@ -7,9 +7,9 @@ import (
|
|||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -50,28 +50,21 @@ func main() {
|
|||||||
os.Exit(0)
|
os.Exit(0)
|
||||||
}
|
}
|
||||||
|
|
||||||
l := logging.NewLogger(os.Stdout)
|
l := logrus.New()
|
||||||
|
l.Out = os.Stdout
|
||||||
|
|
||||||
if *serviceFlag != "" {
|
if *serviceFlag != "" {
|
||||||
if *configTest {
|
if err := doService(configPath, configTest, Build, serviceFlag); err != nil {
|
||||||
fmt.Println("-test is not supported with -service, run the config test without -service")
|
l.WithError(err).Error("Service command failed")
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := doService(configPath, Build, serviceFlag); err != nil {
|
|
||||||
l.Error("Service command failed", "error", err)
|
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
p, err := config.DefaultPath()
|
fmt.Println("-config flag must be set")
|
||||||
if err != nil {
|
flag.Usage()
|
||||||
fmt.Println(err)
|
os.Exit(1)
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
*configPath = p
|
|
||||||
}
|
}
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
@@ -81,16 +74,6 @@ func main() {
|
|||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
|
||||||
fmt.Printf("failed to apply logging config: %s", err)
|
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
|
||||||
l.Error("Failed to reconfigure logger on reload", "error", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
ctrl, err := nebula.Main(c, *configTest, Build, l, nil)
|
ctrl, err := nebula.Main(c, *configTest, Build, l, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
util.LogWithContextIfNeeded("Failed to start", err, l)
|
util.LogWithContextIfNeeded("Failed to start", err, l)
|
||||||
@@ -98,15 +81,16 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !*configTest {
|
if !*configTest {
|
||||||
if err := ctrl.Start(); err != nil {
|
wait, err := ctrl.Start()
|
||||||
|
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 := ctrl.Wait(); err != nil {
|
if err := wait(); err != nil {
|
||||||
l.Error("Nebula stopped due to fatal error", "error", err)
|
l.WithError(err).Error("Nebula stopped due to fatal error")
|
||||||
os.Exit(2)
|
os.Exit(2)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -4,17 +4,19 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
"os"
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
|
||||||
"github.com/kardianos/service"
|
"github.com/kardianos/service"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
var logger service.Logger
|
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
|
||||||
}
|
}
|
||||||
@@ -23,7 +25,8 @@ func (p *program) Start(s service.Service) error {
|
|||||||
// Start should not block.
|
// Start should not block.
|
||||||
logger.Info("Nebula service starting.")
|
logger.Info("Nebula service starting.")
|
||||||
|
|
||||||
l := newPlatformLogger()
|
l := logrus.New()
|
||||||
|
HookLogger(l)
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
err := c.Load(*p.configPath)
|
err := c.Load(*p.configPath)
|
||||||
@@ -31,56 +34,39 @@ func (p *program) Start(s service.Service) error {
|
|||||||
return fmt.Errorf("failed to load config: %s", err)
|
return fmt.Errorf("failed to load config: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
p.control, err = nebula.Main(c, *p.configTest, Build, l, nil)
|
||||||
return fmt.Errorf("failed to apply logging config: %s", err)
|
|
||||||
}
|
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
|
||||||
l.Error("Failed to reconfigure logger on reload", "error", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
p.control, err = nebula.Main(c, false, Build, l, nil)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := p.control.Start(); err != nil {
|
p.control.Start()
|
||||||
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 doService(configPath *string, build string, serviceFlag *string) error {
|
func fileExists(filename string) bool {
|
||||||
|
_, 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 == "" {
|
||||||
p, err := config.DefaultPath()
|
ex, err := os.Executable()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
*configPath = p
|
*configPath = filepath.Dir(ex) + "/config.yaml"
|
||||||
|
if !fileExists(*configPath) {
|
||||||
|
*configPath = filepath.Dir(ex) + "/config.yml"
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
svcConfig := &service.Config{
|
svcConfig := &service.Config{
|
||||||
@@ -92,13 +78,14 @@ func doService(configPath *string, build string, serviceFlag *string) error {
|
|||||||
|
|
||||||
prg := &program{
|
prg := &program{
|
||||||
configPath: configPath,
|
configPath: configPath,
|
||||||
|
configTest: configTest,
|
||||||
build: build,
|
build: build,
|
||||||
}
|
}
|
||||||
|
|
||||||
// Here are what the different loggers are doing:
|
// Here are what the different loggers are doing:
|
||||||
// - `log` is the standard go log utility, meant to be used while the process is still attached to stdout/stderr
|
// - `log` is the standard go log utility, meant to be used while the process is still attached to stdout/stderr
|
||||||
// - `logger` is the service log utility that may be attached to a special place depending on OS (Windows will have it attached to the event log)
|
// - `logger` is the service log utility that may be attached to a special place depending on OS (Windows will have it attached to the event log)
|
||||||
// - in program.Start we build a *slog.Logger via newPlatformLogger; on non-Windows that is a stdout-backed slog logger, on Windows it routes records through the service logger
|
// - above, in `Run` we create a `logrus.Logger` which is what nebula expects to use
|
||||||
s, err := service.New(prg, svcConfig)
|
s, err := service.New(prg, svcConfig)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -123,9 +110,8 @@ func doService(configPath *string, build string, serviceFlag *string) error {
|
|||||||
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 and report the failure
|
// Route any errors to the system logger
|
||||||
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 {
|
||||||
|
|||||||
@@ -1,96 +0,0 @@
|
|||||||
//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])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,5 +0,0 @@
|
|||||||
//go:build fips140enforce
|
|
||||||
|
|
||||||
//go:debug fips140=only
|
|
||||||
|
|
||||||
package main
|
|
||||||
+10
-21
@@ -7,9 +7,9 @@ import (
|
|||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -50,15 +50,13 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
p, err := config.DefaultPath()
|
fmt.Println("-config flag must be set")
|
||||||
if err != nil {
|
flag.Usage()
|
||||||
fmt.Println(err)
|
os.Exit(1)
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
*configPath = p
|
|
||||||
}
|
}
|
||||||
|
|
||||||
l := logging.NewLogger(os.Stdout)
|
l := logrus.New()
|
||||||
|
l.Out = os.Stdout
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
err := c.Load(*configPath)
|
err := c.Load(*configPath)
|
||||||
@@ -67,16 +65,6 @@ func main() {
|
|||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
|
||||||
fmt.Printf("failed to apply logging config: %s", err)
|
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
|
||||||
l.Error("Failed to reconfigure logger on reload", "error", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
ctrl, err := nebula.Main(c, *configTest, Build, l, nil)
|
ctrl, err := nebula.Main(c, *configTest, Build, l, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
util.LogWithContextIfNeeded("Failed to start", err, l)
|
util.LogWithContextIfNeeded("Failed to start", err, l)
|
||||||
@@ -84,7 +72,8 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !*configTest {
|
if !*configTest {
|
||||||
if err := ctrl.Start(); err != nil {
|
wait, err := ctrl.Start()
|
||||||
|
if err != nil {
|
||||||
util.LogWithContextIfNeeded("Error while running", err, l)
|
util.LogWithContextIfNeeded("Error while running", err, l)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
@@ -92,8 +81,8 @@ func main() {
|
|||||||
go ctrl.ShutdownBlock()
|
go ctrl.ShutdownBlock()
|
||||||
notifyReady(l)
|
notifyReady(l)
|
||||||
|
|
||||||
if err := ctrl.Wait(); err != nil {
|
if err := wait(); err != nil {
|
||||||
l.Error("Nebula stopped due to fatal error", "error", err)
|
l.WithError(err).Error("Nebula stopped due to fatal error")
|
||||||
os.Exit(2)
|
os.Exit(2)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,10 +1,11 @@
|
|||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"log/slog"
|
|
||||||
"net"
|
"net"
|
||||||
"os"
|
"os"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
)
|
)
|
||||||
|
|
||||||
// SdNotifyReady tells systemd the service is ready and dependent services can now be started
|
// SdNotifyReady tells systemd the service is ready and dependent services can now be started
|
||||||
@@ -12,30 +13,30 @@ import (
|
|||||||
// https://www.freedesktop.org/software/systemd/man/systemd.service.html
|
// https://www.freedesktop.org/software/systemd/man/systemd.service.html
|
||||||
const SdNotifyReady = "READY=1"
|
const SdNotifyReady = "READY=1"
|
||||||
|
|
||||||
func notifyReady(l *slog.Logger) {
|
func notifyReady(l *logrus.Logger) {
|
||||||
sockName := os.Getenv("NOTIFY_SOCKET")
|
sockName := os.Getenv("NOTIFY_SOCKET")
|
||||||
if sockName == "" {
|
if sockName == "" {
|
||||||
l.Debug("NOTIFY_SOCKET systemd env var not set, not sending ready signal")
|
l.Debugln("NOTIFY_SOCKET systemd env var not set, not sending ready signal")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
conn, err := net.DialTimeout("unixgram", sockName, time.Second)
|
conn, err := net.DialTimeout("unixgram", sockName, time.Second)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
l.Error("failed to connect to systemd notification socket", "error", err)
|
l.WithError(err).Error("failed to connect to systemd notification socket")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
|
|
||||||
err = conn.SetWriteDeadline(time.Now().Add(time.Second))
|
err = conn.SetWriteDeadline(time.Now().Add(time.Second))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
l.Error("failed to set the write deadline for the systemd notification socket", "error", err)
|
l.WithError(err).Error("failed to set the write deadline for the systemd notification socket")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err = conn.Write([]byte(SdNotifyReady)); err != nil {
|
if _, err = conn.Write([]byte(SdNotifyReady)); err != nil {
|
||||||
l.Error("failed to signal the systemd notification socket", "error", err)
|
l.WithError(err).Error("failed to signal the systemd notification socket")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
l.Debug("notified systemd the service is ready")
|
l.Debugln("notified systemd the service is ready")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,8 +3,8 @@
|
|||||||
|
|
||||||
package main
|
package main
|
||||||
|
|
||||||
import "log/slog"
|
import "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
func notifyReady(_ *slog.Logger) {
|
func notifyReady(_ *logrus.Logger) {
|
||||||
// No init service to notify
|
// No init service to notify
|
||||||
}
|
}
|
||||||
|
|||||||
+6
-15
@@ -4,7 +4,6 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
|
||||||
"math"
|
"math"
|
||||||
"os"
|
"os"
|
||||||
"os/signal"
|
"os/signal"
|
||||||
@@ -17,6 +16,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"dario.cat/mergo"
|
"dario.cat/mergo"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"go.yaml.in/yaml/v3"
|
"go.yaml.in/yaml/v3"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -26,11 +26,11 @@ type C struct {
|
|||||||
Settings map[string]any
|
Settings map[string]any
|
||||||
oldSettings map[string]any
|
oldSettings map[string]any
|
||||||
callbacks []func(*C)
|
callbacks []func(*C)
|
||||||
l *slog.Logger
|
l *logrus.Logger
|
||||||
reloadLock sync.Mutex
|
reloadLock sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewC(l *slog.Logger) *C {
|
func NewC(l *logrus.Logger) *C {
|
||||||
return &C{
|
return &C{
|
||||||
Settings: make(map[string]any),
|
Settings: make(map[string]any),
|
||||||
l: l,
|
l: l,
|
||||||
@@ -107,18 +107,12 @@ func (c *C) HasChanged(k string) bool {
|
|||||||
|
|
||||||
newVals, err := yaml.Marshal(nv)
|
newVals, err := yaml.Marshal(nv)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.l.Error("Error while marshaling new config",
|
c.l.WithField("config_path", k).WithError(err).Error("Error while marshaling new config")
|
||||||
"config_path", k,
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
oldVals, err := yaml.Marshal(ov)
|
oldVals, err := yaml.Marshal(ov)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.l.Error("Error while marshaling old config",
|
c.l.WithField("config_path", k).WithError(err).Error("Error while marshaling old config")
|
||||||
"config_path", k,
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return string(newVals) != string(oldVals)
|
return string(newVals) != string(oldVals)
|
||||||
@@ -160,10 +154,7 @@ func (c *C) ReloadConfig() {
|
|||||||
|
|
||||||
err := c.Load(c.path)
|
err := c.Load(c.path)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.l.Error("Error occurred while reloading config",
|
c.l.WithField("config_path", c.path).WithError(err).Error("Error occurred while reloading config")
|
||||||
"config_path", c.path,
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,29 +0,0 @@
|
|||||||
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)
|
|
||||||
}
|
|
||||||
@@ -1,67 +0,0 @@
|
|||||||
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))
|
|
||||||
}
|
|
||||||
+107
-77
@@ -5,12 +5,13 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
@@ -44,16 +45,19 @@ type connectionManager struct {
|
|||||||
inactivityTimeout atomic.Int64
|
inactivityTimeout atomic.Int64
|
||||||
dropInactive atomic.Bool
|
dropInactive atomic.Bool
|
||||||
|
|
||||||
l *slog.Logger
|
metricsTxPunchy metrics.Counter
|
||||||
|
|
||||||
|
l *logrus.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
func newConnectionManagerFromConfig(l *slog.Logger, c *config.C, hm *HostMap, p *Punchy) *connectionManager {
|
func newConnectionManagerFromConfig(l *logrus.Logger, c *config.C, hm *HostMap, p *Punchy) *connectionManager {
|
||||||
cm := &connectionManager{
|
cm := &connectionManager{
|
||||||
hostMap: hm,
|
hostMap: hm,
|
||||||
l: l,
|
l: l,
|
||||||
punchy: p,
|
punchy: p,
|
||||||
relayUsed: make(map[uint32]struct{}),
|
relayUsed: make(map[uint32]struct{}),
|
||||||
relayUsedLock: &sync.RWMutex{},
|
relayUsedLock: &sync.RWMutex{},
|
||||||
|
metricsTxPunchy: metrics.GetOrRegisterCounter("messages.tx.punchy", nil),
|
||||||
}
|
}
|
||||||
|
|
||||||
cm.reload(c, true)
|
cm.reload(c, true)
|
||||||
@@ -81,10 +85,9 @@ func (cm *connectionManager) reload(c *config.C, initial bool) {
|
|||||||
old := cm.getInactivityTimeout()
|
old := cm.getInactivityTimeout()
|
||||||
cm.inactivityTimeout.Store((int64)(c.GetDuration("tunnels.inactivity_timeout", 10*time.Minute)))
|
cm.inactivityTimeout.Store((int64)(c.GetDuration("tunnels.inactivity_timeout", 10*time.Minute)))
|
||||||
if !initial {
|
if !initial {
|
||||||
cm.l.Info("Inactivity timeout has changed",
|
cm.l.WithField("oldDuration", old).
|
||||||
"oldDuration", old,
|
WithField("newDuration", cm.getInactivityTimeout()).
|
||||||
"newDuration", cm.getInactivityTimeout(),
|
Info("Inactivity timeout has changed")
|
||||||
)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -92,10 +95,9 @@ func (cm *connectionManager) reload(c *config.C, initial bool) {
|
|||||||
old := cm.dropInactive.Load()
|
old := cm.dropInactive.Load()
|
||||||
cm.dropInactive.Store(c.GetBool("tunnels.drop_inactive", false))
|
cm.dropInactive.Store(c.GetBool("tunnels.drop_inactive", false))
|
||||||
if !initial {
|
if !initial {
|
||||||
cm.l.Info("Drop inactive setting has changed",
|
cm.l.WithField("oldBool", old).
|
||||||
"oldBool", old,
|
WithField("newBool", cm.dropInactive.Load()).
|
||||||
"newBool", cm.dropInactive.Load(),
|
Info("Drop inactive setting has changed")
|
||||||
)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -136,6 +138,14 @@ 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()
|
||||||
@@ -246,7 +256,7 @@ func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo
|
|||||||
var err error
|
var err error
|
||||||
index, err = AddRelay(cm.l, newhostinfo, cm.hostMap, r.PeerAddr, nil, r.Type, Requested)
|
index, err = AddRelay(cm.l, newhostinfo, cm.hostMap, r.PeerAddr, nil, r.Type, Requested)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
cm.l.Error("failed to migrate relay to new hostinfo", "error", err)
|
cm.l.WithError(err).Error("failed to migrate relay to new hostinfo")
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
switch r.Type {
|
switch r.Type {
|
||||||
@@ -294,16 +304,16 @@ func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo
|
|||||||
|
|
||||||
msg, err := req.Marshal()
|
msg, err := req.Marshal()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
cm.l.Error("failed to marshal Control message to migrate relay", "error", err)
|
cm.l.WithError(err).Error("failed to marshal Control message to migrate relay")
|
||||||
} 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.WithFields(logrus.Fields{
|
||||||
"relayFrom", relayFrom,
|
"relayFrom": req.RelayFromAddr,
|
||||||
"relayTo", relayTo,
|
"relayTo": req.RelayToAddr,
|
||||||
"initiatorRelayIndex", req.InitiatorRelayIndex,
|
"initiatorRelayIndex": req.InitiatorRelayIndex,
|
||||||
"responderRelayIndex", req.ResponderRelayIndex,
|
"responderRelayIndex": req.ResponderRelayIndex,
|
||||||
"vpnAddrs", newhostinfo.vpnAddrs,
|
"vpnAddrs": newhostinfo.vpnAddrs}).
|
||||||
)
|
Info("send CreateRelayRequest")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -315,7 +325,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
|
|
||||||
hostinfo := cm.hostMap.Indexes[localIndex]
|
hostinfo := cm.hostMap.Indexes[localIndex]
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
cm.l.Debug("Not found in hostmap", "localIndex", localIndex)
|
cm.l.WithField("localIndex", localIndex).Debugln("Not found in hostmap")
|
||||||
return doNothing, nil, nil
|
return doNothing, nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -335,10 +345,10 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
// A hostinfo is determined alive if there is incoming traffic
|
// A hostinfo is determined alive if there is incoming traffic
|
||||||
if inTraffic {
|
if inTraffic {
|
||||||
decision := doNothing
|
decision := doNothing
|
||||||
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if cm.l.Level >= logrus.DebugLevel {
|
||||||
hostinfo.logger(cm.l).Debug("Tunnel status",
|
hostinfo.logger(cm.l).
|
||||||
"tunnelCheck", m{"state": "alive", "method": "passive"},
|
WithField("tunnelCheck", m{"state": "alive", "method": "passive"}).
|
||||||
)
|
Debug("Tunnel status")
|
||||||
}
|
}
|
||||||
hostinfo.pendingDeletion.Store(false)
|
hostinfo.pendingDeletion.Store(false)
|
||||||
|
|
||||||
@@ -357,7 +367,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
|
|
||||||
if !outTraffic {
|
if !outTraffic {
|
||||||
// Send a punch packet to keep the NAT state alive
|
// Send a punch packet to keep the NAT state alive
|
||||||
cm.punchy.SendPunch(hostinfo)
|
cm.sendPunch(hostinfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
return decision, hostinfo, primary
|
return decision, hostinfo, primary
|
||||||
@@ -365,9 +375,9 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
|
|
||||||
if hostinfo.pendingDeletion.Load() {
|
if hostinfo.pendingDeletion.Load() {
|
||||||
// We have already sent a test packet and nothing was returned, this hostinfo is dead
|
// We have already sent a test packet and nothing was returned, this hostinfo is dead
|
||||||
hostinfo.logger(cm.l).Info("Tunnel status",
|
hostinfo.logger(cm.l).
|
||||||
"tunnelCheck", m{"state": "dead", "method": "active"},
|
WithField("tunnelCheck", m{"state": "dead", "method": "active"}).
|
||||||
)
|
Info("Tunnel status")
|
||||||
|
|
||||||
return deleteTunnel, hostinfo, nil
|
return deleteTunnel, hostinfo, nil
|
||||||
}
|
}
|
||||||
@@ -378,39 +388,40 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
inactiveFor, isInactive := cm.isInactive(hostinfo, now)
|
inactiveFor, isInactive := cm.isInactive(hostinfo, now)
|
||||||
if isInactive {
|
if isInactive {
|
||||||
// Tunnel is inactive, tear it down
|
// Tunnel is inactive, tear it down
|
||||||
hostinfo.logger(cm.l).Info("Dropping tunnel due to inactivity",
|
hostinfo.logger(cm.l).
|
||||||
"inactiveDuration", inactiveFor,
|
WithField("inactiveDuration", inactiveFor).
|
||||||
"primary", mainHostInfo,
|
WithField("primary", mainHostInfo).
|
||||||
)
|
Info("Dropping tunnel due to inactivity")
|
||||||
|
|
||||||
return closeTunnel, hostinfo, primary
|
return closeTunnel, hostinfo, primary
|
||||||
}
|
}
|
||||||
|
|
||||||
// If we aren't sending or receiving traffic then its an unused tunnel and we don't to test the tunnel.
|
// If we aren't sending or receiving traffic then its an unused tunnel and we don't to test the tunnel.
|
||||||
// Just maintain NAT state if configured to do so.
|
// Just maintain NAT state if configured to do so.
|
||||||
cm.punchy.SendPunch(hostinfo)
|
cm.sendPunch(hostinfo)
|
||||||
cm.trafficTimer.Add(hostinfo.localIndexId, cm.checkInterval)
|
cm.trafficTimer.Add(hostinfo.localIndexId, cm.checkInterval)
|
||||||
return doNothing, nil, nil
|
return doNothing, nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// We aren't receiving traffic but we are sending it. The outbound
|
if cm.punchy.GetTargetEverything() {
|
||||||
// traffic itself refreshes the primary remote's NAT state; this
|
// This is similar to the old punchy behavior with a slight optimization.
|
||||||
// fans out to non-primary remotes, but only if target_all_remotes
|
// We aren't receiving traffic but we are sending it, punch on all known
|
||||||
// is configured.
|
// ips in case we need to re-prime NAT state
|
||||||
cm.punchy.SendPunchToAll(hostinfo)
|
cm.sendPunch(hostinfo)
|
||||||
|
}
|
||||||
|
|
||||||
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if cm.l.Level >= logrus.DebugLevel {
|
||||||
hostinfo.logger(cm.l).Debug("Tunnel status",
|
hostinfo.logger(cm.l).
|
||||||
"tunnelCheck", m{"state": "testing", "method": "active"},
|
WithField("tunnelCheck", m{"state": "testing", "method": "active"}).
|
||||||
)
|
Debug("Tunnel status")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Send a test packet to trigger an authenticated tunnel test, this should suss out any lingering tunnel issues
|
// Send a test packet to trigger an authenticated tunnel test, this should suss out any lingering tunnel issues
|
||||||
decision = sendTestPacket
|
decision = sendTestPacket
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if cm.l.Level >= logrus.DebugLevel {
|
||||||
hostinfo.logger(cm.l).Debug("Hostinfo sadness")
|
hostinfo.logger(cm.l).Debugf("Hostinfo sadness")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -482,16 +493,14 @@ func (cm *connectionManager) isInvalidCertificate(now time.Time, hostinfo *HostI
|
|||||||
return false //cert is still valid! yay!
|
return false //cert is still valid! yay!
|
||||||
} else if err == cert.ErrBlockListed { //avoiding errors.Is for speed
|
} else if err == cert.ErrBlockListed { //avoiding errors.Is for speed
|
||||||
// Block listed certificates should always be disconnected
|
// Block listed certificates should always be disconnected
|
||||||
hostinfo.logger(cm.l).Info("Remote certificate is blocked, tearing down the tunnel",
|
hostinfo.logger(cm.l).WithError(err).
|
||||||
"error", err,
|
WithField("fingerprint", remoteCert.Fingerprint).
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
Info("Remote certificate is blocked, tearing down the tunnel")
|
||||||
)
|
|
||||||
return true
|
return true
|
||||||
} else if cm.intf.disconnectInvalid.Load() {
|
} else if cm.intf.disconnectInvalid.Load() {
|
||||||
hostinfo.logger(cm.l).Info("Remote certificate is no longer valid, tearing down the tunnel",
|
hostinfo.logger(cm.l).WithError(err).
|
||||||
"error", err,
|
WithField("fingerprint", remoteCert.Fingerprint).
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
Info("Remote certificate is no longer valid, tearing down the tunnel")
|
||||||
)
|
|
||||||
return true
|
return true
|
||||||
} else {
|
} else {
|
||||||
//if we reach here, the cert is no longer valid, but we're configured to keep tunnels from now-invalid certs open
|
//if we reach here, the cert is no longer valid, but we're configured to keep tunnels from now-invalid certs open
|
||||||
@@ -499,17 +508,41 @@ 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
|
||||||
curCrtVersion := curCrt.Version()
|
curCrtVersion := curCrt.Version()
|
||||||
myCrt := cs.getCertificate(curCrtVersion)
|
myCrt := cs.getCertificate(curCrtVersion)
|
||||||
if myCrt == nil {
|
if myCrt == nil {
|
||||||
cm.l.Info("Re-handshaking with remote",
|
cm.l.WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
WithField("version", curCrtVersion).
|
||||||
"version", curCrtVersion,
|
WithField("reason", "local certificate removed").
|
||||||
"reason", "local certificate removed",
|
Info("Re-handshaking with remote")
|
||||||
)
|
|
||||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -517,12 +550,11 @@ func (cm *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
|||||||
if peerCrt != nil && curCrtVersion < peerCrt.Certificate.Version() {
|
if peerCrt != nil && curCrtVersion < peerCrt.Certificate.Version() {
|
||||||
// if our certificate version is less than theirs, and we have a matching version available, rehandshake?
|
// if our certificate version is less than theirs, and we have a matching version available, rehandshake?
|
||||||
if cs.getCertificate(peerCrt.Certificate.Version()) != nil {
|
if cs.getCertificate(peerCrt.Certificate.Version()) != nil {
|
||||||
cm.l.Info("Re-handshaking with remote",
|
cm.l.WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
WithField("version", curCrtVersion).
|
||||||
"version", curCrtVersion,
|
WithField("peerVersion", peerCrt.Certificate.Version()).
|
||||||
"peerVersion", peerCrt.Certificate.Version(),
|
WithField("reason", "local certificate version lower than peer, attempting to correct").
|
||||||
"reason", "local certificate version lower than peer, attempting to correct",
|
Info("Re-handshaking with remote")
|
||||||
)
|
|
||||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], func(hh *HandshakeHostInfo) {
|
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], func(hh *HandshakeHostInfo) {
|
||||||
hh.initiatingVersionOverride = peerCrt.Certificate.Version()
|
hh.initiatingVersionOverride = peerCrt.Certificate.Version()
|
||||||
})
|
})
|
||||||
@@ -530,19 +562,17 @@ func (cm *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if !bytes.Equal(curCrt.Signature(), myCrt.Signature()) {
|
if !bytes.Equal(curCrt.Signature(), myCrt.Signature()) {
|
||||||
cm.l.Info("Re-handshaking with remote",
|
cm.l.WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
WithField("reason", "local certificate is not current").
|
||||||
"reason", "local certificate is not current",
|
Info("Re-handshaking with remote")
|
||||||
)
|
|
||||||
|
|
||||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if curCrtVersion < cs.initiatingVersion {
|
if curCrtVersion < cs.initiatingVersion {
|
||||||
cm.l.Info("Re-handshaking with remote",
|
cm.l.WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
WithField("reason", "current cert version < pki.initiatingVersion").
|
||||||
"reason", "current cert version < pki.initiatingVersion",
|
Info("Re-handshaking with remote")
|
||||||
)
|
|
||||||
|
|
||||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
||||||
return
|
return
|
||||||
|
|||||||
+27
-24
@@ -7,9 +7,9 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/overlay/overlaytest"
|
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
@@ -25,7 +25,6 @@ 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)
|
||||||
|
|
||||||
@@ -47,13 +46,13 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
inside: &overlaytest.NoopTun{},
|
inside: &test.NoopTun{},
|
||||||
outside: &udp.NoopConn{},
|
outside: &udp.NoopConn{},
|
||||||
firewall: &Firewall{},
|
firewall: &Firewall{},
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
@@ -64,9 +63,9 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
ifce.pki.cs.Store(cs)
|
ifce.pki.cs.Store(cs)
|
||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(l)
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(l, conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(l, conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
@@ -80,6 +79,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -129,13 +129,13 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
inside: &overlaytest.NoopTun{},
|
inside: &test.NoopTun{},
|
||||||
outside: &udp.NoopConn{},
|
outside: &udp.NoopConn{},
|
||||||
firewall: &Firewall{},
|
firewall: &Firewall{},
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
@@ -146,9 +146,9 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
ifce.pki.cs.Store(cs)
|
ifce.pki.cs.Store(cs)
|
||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(l)
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(l, conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(l, conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
@@ -162,6 +162,7 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -213,13 +214,13 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
inside: &overlaytest.NoopTun{},
|
inside: &test.NoopTun{},
|
||||||
outside: &udp.NoopConn{},
|
outside: &udp.NoopConn{},
|
||||||
firewall: &Firewall{},
|
firewall: &Firewall{},
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
@@ -230,12 +231,12 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
ifce.pki.cs.Store(cs)
|
ifce.pki.cs.Store(cs)
|
||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(l)
|
||||||
conf.Settings["tunnels"] = map[string]any{
|
conf.Settings["tunnels"] = map[string]any{
|
||||||
"drop_inactive": true,
|
"drop_inactive": true,
|
||||||
}
|
}
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(l, conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(l, conf, hostMap, punchy)
|
||||||
assert.True(t, nc.dropInactive.Load())
|
assert.True(t, nc.dropInactive.Load())
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
|
|
||||||
@@ -247,6 +248,7 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -337,15 +339,15 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
cachedPeerCert, err := ncp.VerifyCertificate(now.Add(time.Second), peerCert)
|
cachedPeerCert, err := ncp.VerifyCertificate(now.Add(time.Second), peerCert)
|
||||||
|
|
||||||
cs := &CertState{
|
cs := &CertState{
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{},
|
v1Cert: &dummyCert{},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
inside: &overlaytest.NoopTun{},
|
inside: &test.NoopTun{},
|
||||||
outside: &udp.NoopConn{},
|
outside: &udp.NoopConn{},
|
||||||
firewall: &Firewall{},
|
firewall: &Firewall{},
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
@@ -358,9 +360,9 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
ifce.disconnectInvalid.Store(true)
|
ifce.disconnectInvalid.Store(true)
|
||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(l)
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(l, conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(l, conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
ifce.connectionManager = nc
|
ifce.connectionManager = nc
|
||||||
|
|
||||||
@@ -369,6 +371,7 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
ConnectionState: &ConnectionState{
|
ConnectionState: &ConnectionState{
|
||||||
myCert: &dummyCert{},
|
myCert: &dummyCert{},
|
||||||
peerCert: cachedPeerCert,
|
peerCert: cachedPeerCert,
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|||||||
+53
-70
@@ -1,49 +1,81 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"crypto/rand"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"log/slog"
|
"fmt"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/handshake"
|
|
||||||
"github.com/slackhq/nebula/header"
|
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
)
|
)
|
||||||
|
|
||||||
const ReplayWindow = 1024
|
const ReplayWindow = 1024
|
||||||
|
|
||||||
type ConnectionState struct {
|
type ConnectionState struct {
|
||||||
eKey noiseutil.CipherState
|
eKey *NebulaCipherState
|
||||||
dKey noiseutil.CipherState
|
dKey *NebulaCipherState
|
||||||
|
H *noise.HandshakeState
|
||||||
myCert cert.Certificate
|
myCert cert.Certificate
|
||||||
peerCert *cert.CachedCertificate
|
peerCert *cert.CachedCertificate
|
||||||
initiator bool
|
initiator bool
|
||||||
messageCounter atomic.Uint64
|
messageCounter atomic.Uint64
|
||||||
window *Bits
|
window *Bits
|
||||||
decryptLock sync.Mutex
|
|
||||||
writeLock sync.Mutex
|
writeLock sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
// newConnectionStateFromResult builds a fully-populated ConnectionState from a
|
func NewConnectionState(l *logrus.Logger, cs *CertState, crt cert.Certificate, initiator bool, pattern noise.HandshakePattern) (*ConnectionState, error) {
|
||||||
// completed handshake.Result. It seeds messageCounter and the replay window so
|
var dhFunc noise.DHFunc
|
||||||
// that the post-handshake message indices already used on the wire don't count
|
switch crt.Curve() {
|
||||||
// as missed traffic in the data plane.
|
case cert.Curve_CURVE25519:
|
||||||
func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
|
dhFunc = noise.DH25519
|
||||||
|
case cert.Curve_P256:
|
||||||
|
if cs.pkcs11Backed {
|
||||||
|
dhFunc = noiseutil.DHP256PKCS11
|
||||||
|
} else {
|
||||||
|
dhFunc = noiseutil.DHP256
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return nil, fmt.Errorf("invalid curve: %s", crt.Curve())
|
||||||
|
}
|
||||||
|
|
||||||
|
var ncs noise.CipherSuite
|
||||||
|
if cs.cipher == "chachapoly" {
|
||||||
|
ncs = noise.NewCipherSuite(dhFunc, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||||
|
} else {
|
||||||
|
ncs = noise.NewCipherSuite(dhFunc, noiseutil.CipherAESGCM, noise.HashSHA256)
|
||||||
|
}
|
||||||
|
|
||||||
|
static := noise.DHKey{Private: cs.privateKey, Public: crt.PublicKey()}
|
||||||
|
hs, err := noise.NewHandshakeState(noise.Config{
|
||||||
|
CipherSuite: ncs,
|
||||||
|
Random: rand.Reader,
|
||||||
|
Pattern: pattern,
|
||||||
|
Initiator: initiator,
|
||||||
|
StaticKeypair: static,
|
||||||
|
//NOTE: These should come from CertState (pki.go) when we finally implement it
|
||||||
|
PresharedKey: []byte{},
|
||||||
|
PresharedKeyPlacement: 0,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("NewConnectionState: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The queue and ready params prevent a counter race that would happen when
|
||||||
|
// sending stored packets and simultaneously accepting new traffic.
|
||||||
ci := &ConnectionState{
|
ci := &ConnectionState{
|
||||||
myCert: r.MyCert,
|
H: hs,
|
||||||
initiator: r.Initiator,
|
initiator: initiator,
|
||||||
peerCert: r.RemoteCert,
|
|
||||||
eKey: noiseutil.NewCipherState(r.EKey, r.Cipher),
|
|
||||||
dKey: noiseutil.NewCipherState(r.DKey, r.Cipher),
|
|
||||||
window: NewBits(ReplayWindow),
|
window: NewBits(ReplayWindow),
|
||||||
|
myCert: crt,
|
||||||
}
|
}
|
||||||
ci.messageCounter.Add(r.MessageIndex)
|
// always start the counter from 2, as packet 1 and packet 2 are handshake packets.
|
||||||
for i := uint64(1); i <= r.MessageIndex; i++ {
|
ci.messageCounter.Add(2)
|
||||||
ci.window.Update(nil, i)
|
|
||||||
}
|
return ci, nil
|
||||||
return ci
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
||||||
@@ -57,52 +89,3 @@ 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
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,114 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
ct "github.com/slackhq/nebula/cert_test"
|
|
||||||
"github.com/slackhq/nebula/handshake"
|
|
||||||
"github.com/slackhq/nebula/header"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
// runTestHandshake runs a complete IX handshake between two freshly-built
|
|
||||||
// peers and returns the initiator and responder Results. Used to produce
|
|
||||||
// real cipher states for tests that need to exercise post-handshake glue.
|
|
||||||
func runTestHandshake(t *testing.T) (initR, respR *handshake.Result) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
|
|
||||||
makeCreds := func(name string, networks []netip.Prefix) handshake.GetCredentialFunc {
|
|
||||||
c, _, rawKey, _ := ct.NewTestCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, ca, caKey,
|
|
||||||
name, ca.NotBefore(), ca.NotAfter(), networks, nil, nil,
|
|
||||||
)
|
|
||||||
priv, _, _, err := cert.UnmarshalPrivateKeyFromPEM(rawKey)
|
|
||||||
require.NoError(t, err)
|
|
||||||
hsBytes, err := c.MarshalForHandshakes()
|
|
||||||
require.NoError(t, err)
|
|
||||||
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
|
||||||
cred := handshake.NewCredential(c, hsBytes, priv, ncs)
|
|
||||||
return func(v cert.Version) *handshake.Credential {
|
|
||||||
if v == cert.Version2 {
|
|
||||||
return cred
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
verifier := func(c cert.Certificate) (*cert.CachedCertificate, error) {
|
|
||||||
return caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
}
|
|
||||||
|
|
||||||
initCreds := makeCreds("initiator", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
respCreds := makeCreds("responder", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
|
||||||
|
|
||||||
initM, err := handshake.NewMachine(
|
|
||||||
cert.Version2, initCreds, verifier,
|
|
||||||
func() (uint32, error) { return 1000, nil },
|
|
||||||
true, header.HandshakeIXPSK0,
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
respM, err := handshake.NewMachine(
|
|
||||||
cert.Version2, respCreds, verifier,
|
|
||||||
func() (uint32, error) { return 2000, nil },
|
|
||||||
false, header.HandshakeIXPSK0,
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
msg1, err := initM.Initiate(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
resp, respR, err := respM.ProcessPacket(nil, msg1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, respR)
|
|
||||||
|
|
||||||
_, initR, err = initM.ProcessPacket(nil, resp)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, initR)
|
|
||||||
|
|
||||||
return initR, respR
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewConnectionStateFromResult(t *testing.T) {
|
|
||||||
initR, respR := runTestHandshake(t)
|
|
||||||
|
|
||||||
t.Run("initiator", func(t *testing.T) {
|
|
||||||
ci := newConnectionStateFromResult(initR)
|
|
||||||
assert.True(t, ci.initiator)
|
|
||||||
assert.Equal(t, initR.MyCert, ci.myCert)
|
|
||||||
assert.Equal(t, initR.RemoteCert, ci.peerCert)
|
|
||||||
assert.NotNil(t, ci.eKey)
|
|
||||||
assert.NotNil(t, ci.dKey)
|
|
||||||
|
|
||||||
// IX has 2 handshake messages; the next data-plane send is counter=3.
|
|
||||||
assert.Equal(t, uint64(2), ci.messageCounter.Load(),
|
|
||||||
"messageCounter must equal Result.MessageIndex so the next send is N+1")
|
|
||||||
|
|
||||||
// Both handshake counters must be marked seen so they don't appear lost.
|
|
||||||
// Check returns false if an index has already been recorded.
|
|
||||||
assert.False(t, ci.window.Check(nil, 1), "counter 1 must already be seen")
|
|
||||||
assert.False(t, ci.window.Check(nil, 2), "counter 2 must already be seen")
|
|
||||||
// Counter 3 is the next data-plane message and must NOT be pre-marked.
|
|
||||||
assert.True(t, ci.window.Check(nil, 3), "counter 3 must not be pre-seeded")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("responder", func(t *testing.T) {
|
|
||||||
ci := newConnectionStateFromResult(respR)
|
|
||||||
assert.False(t, ci.initiator)
|
|
||||||
assert.Equal(t, respR.MyCert, ci.myCert)
|
|
||||||
assert.Equal(t, respR.RemoteCert, ci.peerCert)
|
|
||||||
assert.NotNil(t, ci.eKey)
|
|
||||||
assert.NotNil(t, ci.dKey)
|
|
||||||
assert.Equal(t, uint64(2), ci.messageCounter.Load())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
+62
-67
@@ -3,14 +3,15 @@ package nebula
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"log/slog"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"os/signal"
|
"os/signal"
|
||||||
"sync"
|
"sync"
|
||||||
"syscall"
|
"syscall"
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/firewall/events"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
)
|
)
|
||||||
@@ -46,14 +47,13 @@ type Control struct {
|
|||||||
state RunState
|
state RunState
|
||||||
|
|
||||||
f *Interface
|
f *Interface
|
||||||
l *slog.Logger
|
l *logrus.Logger
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
sshStart func()
|
sshStart func()
|
||||||
statsStart func()
|
statsStart func()
|
||||||
dnsStart func()
|
dnsStart func()
|
||||||
lighthouseStart func()
|
lighthouseStart func()
|
||||||
networkChangeStart func(rebind func())
|
|
||||||
connectionManagerStart func(context.Context)
|
connectionManagerStart func(context.Context)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -70,29 +70,29 @@ type ControlHostInfo struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Start actually runs nebula, this is a nonblocking call.
|
// Start actually runs nebula, this is a nonblocking call.
|
||||||
// Use Wait to block until nebula has fully stopped and to learn whether a fatal reader error caused the shutdown.
|
// The returned function blocks until nebula has fully stopped and returns the
|
||||||
func (c *Control) Start() error {
|
// first fatal reader error (if any). A nil error means nebula shut down
|
||||||
|
// gracefully; a non-nil error means a reader hit an unexpected failure that
|
||||||
|
// triggered the shutdown.
|
||||||
|
func (c *Control) Start() (func() error, error) {
|
||||||
c.stateLock.Lock()
|
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 ErrAlreadyStopped
|
return nil, ErrAlreadyStopped
|
||||||
case StateStarted:
|
case StateStarted:
|
||||||
return ErrAlreadyStarted
|
return nil, ErrAlreadyStarted
|
||||||
default:
|
default:
|
||||||
return ErrUnknownState
|
return nil, 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 err
|
return nil, 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.
|
||||||
@@ -105,9 +105,6 @@ func (c *Control) Start() 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)
|
||||||
}
|
}
|
||||||
@@ -118,9 +115,13 @@ func (c *Control) Start() error {
|
|||||||
c.f.triggerShutdown = c.Stop
|
c.f.triggerShutdown = c.Stop
|
||||||
|
|
||||||
// Start reading packets.
|
// Start reading packets.
|
||||||
c.f.run()
|
out, err := c.f.run()
|
||||||
|
if err != nil {
|
||||||
|
c.state = StateStopped
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
c.state = StateStarted
|
c.state = StateStarted
|
||||||
return nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) State() RunState {
|
func (c *Control) State() RunState {
|
||||||
@@ -133,26 +134,10 @@ func (c *Control) Context() context.Context {
|
|||||||
return c.ctx
|
return c.ctx
|
||||||
}
|
}
|
||||||
|
|
||||||
// Stop tears nebula down, closing all tunnels and releasing everything it holds.
|
// Stop is a non-blocking call that signals nebula to close all tunnels and shut down
|
||||||
// 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()
|
||||||
switch c.state {
|
if c.state != StateStarted {
|
||||||
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
|
||||||
@@ -161,26 +146,19 @@ func (c *Control) Stop() {
|
|||||||
c.state = StateStopping
|
c.state = StateStopping
|
||||||
c.stateLock.Unlock()
|
c.stateLock.Unlock()
|
||||||
|
|
||||||
// Closing tunnels can be slow with a large hostmap, don't hold the lock for it
|
// Stop the handshakeManager (and other services), to prevent new tunnels from
|
||||||
|
// being created while we're shutting them all down.
|
||||||
c.cancel()
|
c.cancel()
|
||||||
c.CloseAllTunnels(false)
|
|
||||||
|
|
||||||
|
c.CloseAllTunnels(false)
|
||||||
|
if err := c.f.Close(); err != nil {
|
||||||
|
c.l.WithError(err).Error("Close interface failed")
|
||||||
|
}
|
||||||
c.stateLock.Lock()
|
c.stateLock.Lock()
|
||||||
c.state = StateStopped
|
c.state = StateStopped
|
||||||
if err := c.f.Close(); err != nil {
|
|
||||||
c.l.Error("Close interface failed", "error", err)
|
|
||||||
}
|
|
||||||
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)
|
||||||
@@ -189,24 +167,13 @@ func (c *Control) ShutdownBlock() {
|
|||||||
|
|
||||||
rawSig := <-sigChan
|
rawSig := <-sigChan
|
||||||
sig := rawSig.String()
|
sig := rawSig.String()
|
||||||
c.l.Info("Caught signal, shutting down", "signal", sig)
|
c.l.WithField("signal", sig).Info("Caught signal, shutting down")
|
||||||
c.Stop()
|
c.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change.
|
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change
|
||||||
func (c *Control) RebindUDPServer() {
|
func (c *Control) RebindUDPServer() {
|
||||||
c.stateLock.Lock()
|
_ = c.f.outside.Rebind()
|
||||||
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()
|
||||||
@@ -337,10 +304,8 @@ func (c *Control) CloseAllTunnels(excludeLighthouses bool) (closed int) {
|
|||||||
c.f.send(header.CloseTunnel, 0, h.ConnectionState, h, []byte{}, make([]byte, 12, 12), make([]byte, mtu))
|
c.f.send(header.CloseTunnel, 0, h.ConnectionState, h, []byte{}, make([]byte, 12, 12), make([]byte, mtu))
|
||||||
c.f.closeTunnel(h)
|
c.f.closeTunnel(h)
|
||||||
|
|
||||||
c.l.Debug("Sending close tunnel message",
|
c.l.WithField("vpnAddrs", h.vpnAddrs).WithField("udpAddr", h.remote).
|
||||||
"vpnAddrs", h.vpnAddrs,
|
Debug("Sending close tunnel message")
|
||||||
"udpAddr", h.GetRemote(),
|
|
||||||
)
|
|
||||||
closed++
|
closed++
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -376,6 +341,36 @@ func (c *Control) Device() overlay.Device {
|
|||||||
return c.f.inside
|
return c.f.inside
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetFirewallEventReporter installs an event reporter on the current firewall.
|
||||||
|
// Passing nil clears any installed reporter. The reporter is carried across
|
||||||
|
// firewall rule reloads. Report* methods are invoked while nebula holds
|
||||||
|
// internal locks and must be non-blocking; in particular they must not call
|
||||||
|
// back into *Control methods that touch the firewall, or deadlock will
|
||||||
|
// result.
|
||||||
|
//
|
||||||
|
// Installation is performed by shallow-copying the current *Firewall,
|
||||||
|
// setting the reporter field on the copy, and swapping the pointer under
|
||||||
|
// the conntrack lock. Every Firewall the data path sees therefore has an
|
||||||
|
// immutable reporter slot, and emit sites can read it without any
|
||||||
|
// synchronization of their own.
|
||||||
|
func (c *Control) SetFirewallEventReporter(r events.Reporter) {
|
||||||
|
old := c.f.firewall
|
||||||
|
if old == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
old.Conntrack.Lock()
|
||||||
|
defer old.Conntrack.Unlock()
|
||||||
|
|
||||||
|
// Re-read under the lock in case a concurrent reload swapped in a new
|
||||||
|
// Firewall between the unlocked load above and here. Both Firewalls share
|
||||||
|
// the same Conntrack pointer in the normal (non-overflow) reload path,
|
||||||
|
// so the lock we hold is the right one for whichever we see now.
|
||||||
|
current := c.f.firewall
|
||||||
|
fw := *current
|
||||||
|
fw.reporter = r
|
||||||
|
c.f.firewall = &fw
|
||||||
|
}
|
||||||
|
|
||||||
func copyHostInfo(h *HostInfo, preferredRanges []netip.Prefix) ControlHostInfo {
|
func copyHostInfo(h *HostInfo, preferredRanges []netip.Prefix) ControlHostInfo {
|
||||||
chi := ControlHostInfo{
|
chi := ControlHostInfo{
|
||||||
VpnAddrs: make([]netip.Addr, len(h.vpnAddrs)),
|
VpnAddrs: make([]netip.Addr, len(h.vpnAddrs)),
|
||||||
@@ -384,7 +379,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.GetRemote(),
|
CurrentRemote: h.remote,
|
||||||
}
|
}
|
||||||
|
|
||||||
for i, a := range h.vpnAddrs {
|
for i, a := range h.vpnAddrs {
|
||||||
|
|||||||
@@ -1,292 +0,0 @@
|
|||||||
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")
|
|
||||||
}
|
|
||||||
+8
-162
@@ -1,17 +1,15 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"log/slog"
|
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"reflect"
|
"reflect"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"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) {
|
||||||
@@ -45,7 +43,8 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
|||||||
assert.True(t, ok)
|
assert.True(t, ok)
|
||||||
|
|
||||||
crt := &dummyCert{}
|
crt := &dummyCert{}
|
||||||
hi := &HostInfo{
|
hm.unlockedAddHostInfo(&HostInfo{
|
||||||
|
remote: remote1,
|
||||||
remotes: remotes,
|
remotes: remotes,
|
||||||
ConnectionState: &ConnectionState{
|
ConnectionState: &ConnectionState{
|
||||||
peerCert: &cert.CachedCertificate{Certificate: crt},
|
peerCert: &cert.CachedCertificate{Certificate: crt},
|
||||||
@@ -58,14 +57,13 @@ 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)
|
||||||
|
|
||||||
hi2 := &HostInfo{
|
hm.unlockedAddHostInfo(&HostInfo{
|
||||||
|
remote: remote1,
|
||||||
remotes: remotes,
|
remotes: remotes,
|
||||||
ConnectionState: &ConnectionState{
|
ConnectionState: &ConnectionState{
|
||||||
peerCert: nil,
|
peerCert: nil,
|
||||||
@@ -78,16 +76,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{})
|
||||||
hi2.remote.Store(&remote1)
|
|
||||||
hm.unlockedAddHostInfo(hi2, &Interface{})
|
|
||||||
|
|
||||||
c := Control{
|
c := Control{
|
||||||
state: StateReady,
|
state: StateReady,
|
||||||
f: &Interface{
|
f: &Interface{
|
||||||
hostMap: hm,
|
hostMap: hm,
|
||||||
},
|
},
|
||||||
l: test.NewLogger(),
|
l: logrus.New(),
|
||||||
}
|
}
|
||||||
|
|
||||||
thi := c.GetHostInfoByVpnAddr(vpnIp, false)
|
thi := c.GetHostInfoByVpnAddr(vpnIp, false)
|
||||||
@@ -124,153 +120,3 @@ 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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
+61
-33
@@ -5,6 +5,8 @@ package nebula
|
|||||||
import (
|
import (
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/google/gopacket"
|
||||||
|
"github.com/google/gopacket/layers"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
@@ -20,9 +22,7 @@ func (c *Control) WaitForType(msgType header.MessageType, subType header.Message
|
|||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
pipeTo.InjectUDPPacket(p)
|
pipeTo.InjectUDPPacket(p)
|
||||||
match := h.Type == msgType && h.Subtype == subType
|
if h.Type == msgType && h.Subtype == subType {
|
||||||
p.Release()
|
|
||||||
if match {
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -38,9 +38,7 @@ func (c *Control) WaitForTypeByIndex(toIndex uint32, msgType header.MessageType,
|
|||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
pipeTo.InjectUDPPacket(p)
|
pipeTo.InjectUDPPacket(p)
|
||||||
match := h.RemoteIndex == toIndex && h.Type == msgType && h.Subtype == subType
|
if h.RemoteIndex == toIndex && h.Type == msgType && h.Subtype == subType {
|
||||||
p.Release()
|
|
||||||
if match {
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -92,15 +90,65 @@ func (c *Control) GetTunTxChan() <-chan []byte {
|
|||||||
return c.f.inside.(*overlay.TestTun).TxPackets
|
return c.f.inside.(*overlay.TestTun).TxPackets
|
||||||
}
|
}
|
||||||
|
|
||||||
// InjectUDPPacket injects a packet into the udp side. We copy internally so the caller keeps ownership of p.
|
// InjectUDPPacket will inject a packet into the udp side of nebula
|
||||||
// The copy comes from the freelist so steady-state alloc is zero.
|
|
||||||
func (c *Control) InjectUDPPacket(p *udp.Packet) {
|
func (c *Control) InjectUDPPacket(p *udp.Packet) {
|
||||||
c.f.outside.(*udp.TesterConn).Send(p.Copy())
|
c.f.outside.(*udp.TesterConn).Send(p)
|
||||||
}
|
}
|
||||||
|
|
||||||
// InjectTunPacket pushes an IP packet onto the tun interface.
|
// InjectTunUDPPacket puts a udp packet on the tun interface. Using UDP here because it's a simpler protocol
|
||||||
func (c *Control) InjectTunPacket(packet []byte) {
|
func (c *Control) InjectTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr netip.Addr, fromPort uint16, data []byte) {
|
||||||
c.f.inside.(*overlay.TestTun).Send(packet)
|
serialize := make([]gopacket.SerializableLayer, 0)
|
||||||
|
var netLayer gopacket.NetworkLayer
|
||||||
|
if toAddr.Is6() {
|
||||||
|
if !fromAddr.Is6() {
|
||||||
|
panic("Cant send ipv6 to ipv4")
|
||||||
|
}
|
||||||
|
ip := &layers.IPv6{
|
||||||
|
Version: 6,
|
||||||
|
NextHeader: layers.IPProtocolUDP,
|
||||||
|
SrcIP: fromAddr.Unmap().AsSlice(),
|
||||||
|
DstIP: toAddr.Unmap().AsSlice(),
|
||||||
|
}
|
||||||
|
serialize = append(serialize, ip)
|
||||||
|
netLayer = ip
|
||||||
|
} else {
|
||||||
|
if !fromAddr.Is4() {
|
||||||
|
panic("Cant send ipv4 to ipv6")
|
||||||
|
}
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Protocol: layers.IPProtocolUDP,
|
||||||
|
SrcIP: fromAddr.Unmap().AsSlice(),
|
||||||
|
DstIP: toAddr.Unmap().AsSlice(),
|
||||||
|
}
|
||||||
|
serialize = append(serialize, ip)
|
||||||
|
netLayer = ip
|
||||||
|
}
|
||||||
|
|
||||||
|
udp := layers.UDP{
|
||||||
|
SrcPort: layers.UDPPort(fromPort),
|
||||||
|
DstPort: layers.UDPPort(toPort),
|
||||||
|
}
|
||||||
|
err := udp.SetNetworkLayerForChecksum(netLayer)
|
||||||
|
if err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
buffer := gopacket.NewSerializeBuffer()
|
||||||
|
opt := gopacket.SerializeOptions{
|
||||||
|
ComputeChecksums: true,
|
||||||
|
FixLengths: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
serialize = append(serialize, &udp, gopacket.Payload(data))
|
||||||
|
err = gopacket.SerializeLayers(buffer, opt, serialize...)
|
||||||
|
if err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
c.f.inside.(*overlay.TestTun).Send(buffer.Bytes())
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) GetVpnAddrs() []netip.Addr {
|
func (c *Control) GetVpnAddrs() []netip.Addr {
|
||||||
@@ -108,19 +156,7 @@ func (c *Control) GetVpnAddrs() []netip.Addr {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) GetUDPAddr() netip.AddrPort {
|
func (c *Control) GetUDPAddr() netip.AddrPort {
|
||||||
return c.f.outside.(*udp.TesterConn).GetAddr()
|
return c.f.outside.(*udp.TesterConn).Addr
|
||||||
}
|
|
||||||
|
|
||||||
// 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 {
|
||||||
@@ -137,14 +173,6 @@ 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 = bit.band(tvbuf:range(0,1):uint(), 0x0F)
|
local nebula_type = bit32.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))
|
||||||
|
|||||||
+46
-143
@@ -3,7 +3,6 @@ package nebula
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"strconv"
|
"strconv"
|
||||||
@@ -11,21 +10,20 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
|
"github.com/gaissmai/bart"
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
type dnsServer struct {
|
type dnsServer struct {
|
||||||
sync.RWMutex
|
sync.RWMutex
|
||||||
l *slog.Logger
|
l *logrus.Logger
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
dnsMap4 map[string]netip.Addr
|
dnsMap4 map[string]netip.Addr
|
||||||
dnsMap6 map[string]netip.Addr
|
dnsMap6 map[string]netip.Addr
|
||||||
hostMap *HostMap
|
hostMap *HostMap
|
||||||
pki *PKI
|
myVpnAddrsTable *bart.Lite
|
||||||
|
|
||||||
// selfHost is the cached FQDN we last seeded for ourselves
|
|
||||||
selfHost string
|
|
||||||
|
|
||||||
mux *dns.ServeMux
|
mux *dns.ServeMux
|
||||||
|
|
||||||
@@ -57,28 +55,27 @@ 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, pki *PKI, hostMap *HostMap, c *config.C) (*dnsServer, error) {
|
func newDnsServerFromConfig(ctx context.Context, l *logrus.Logger, cs *CertState, 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,
|
||||||
pki: pki,
|
myVpnAddrsTable: cs.myVpnAddrsTable,
|
||||||
}
|
}
|
||||||
ds.mux = dns.NewServeMux()
|
ds.mux = dns.NewServeMux()
|
||||||
ds.mux.HandleFunc(".", ds.handleDnsRequest)
|
ds.mux.HandleFunc(".", ds.handleDnsRequest)
|
||||||
|
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
c.RegisterReloadCallback(func(c *config.C) {
|
||||||
if err := ds.reload(c, false); err != nil {
|
if err := ds.reload(c, false); err != nil {
|
||||||
ds.l.Error("Failed to reload DNS responder from config", "error", err)
|
l.WithError(err).Error("Failed to reload DNS responder from config")
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -97,7 +94,8 @@ 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 != nil
|
running := d.server
|
||||||
|
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)
|
||||||
@@ -111,26 +109,29 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !enabled {
|
if !enabled {
|
||||||
if running {
|
if running != nil {
|
||||||
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 and a fresh seedSelf.
|
// will repopulate from fresh handshakes.
|
||||||
d.clearRecords()
|
d.clearRecords()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if !running {
|
if running == nil {
|
||||||
// Was disabled (or never started); bring it up now.
|
// Was disabled (or never started); bring it up now.
|
||||||
go d.Start()
|
go d.Start()
|
||||||
} else if !sameAddr {
|
return nil
|
||||||
// Stop clears the slot before shutting down, otherwise the Start below can find the dying server and refuse
|
|
||||||
d.Stop()
|
|
||||||
go d.Start()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Refresh the self entry every enabled reload so cert renewals that change our name or VPN addresses are picked up.
|
if sameAddr {
|
||||||
d.seedSelf()
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
d.shutdownServer(running, runningStarted, "reload")
|
||||||
|
// Old Start goroutine has now exited; bring up a fresh listener on the
|
||||||
|
// new address.
|
||||||
|
go d.Start()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -144,7 +145,7 @@ func (d *dnsServer) shutdownServer(srv *dns.Server, started chan struct{}, reaso
|
|||||||
<-started
|
<-started
|
||||||
}
|
}
|
||||||
if err := srv.Shutdown(); err != nil {
|
if err := srv.Shutdown(); err != nil {
|
||||||
d.l.Warn("Failed to shut down the DNS responder", "reason", reason, "error", err)
|
d.l.WithError(err).WithField("reason", reason).Warn("Failed to shut down the DNS responder")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -161,9 +162,7 @@ func (d *dnsServer) Start() {
|
|||||||
|
|
||||||
started := make(chan struct{})
|
started := make(chan struct{})
|
||||||
d.serverMu.Lock()
|
d.serverMu.Lock()
|
||||||
// Re-check enabled under the lock, a disable that raced our check above snapshots the slot under it too.
|
if d.ctx.Err() != nil {
|
||||||
// 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
|
||||||
}
|
}
|
||||||
@@ -189,7 +188,7 @@ func (d *dnsServer) Start() {
|
|||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
d.l.Info("Starting DNS responder", "dnsListener", addr)
|
d.l.WithField("dnsListener", addr).Info("Starting DNS responder")
|
||||||
err := server.ListenAndServe()
|
err := server.ListenAndServe()
|
||||||
close(done)
|
close(done)
|
||||||
|
|
||||||
@@ -201,16 +200,8 @@ 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.WithError(err).Warn("Failed to run the DNS responder")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -225,53 +216,30 @@ func (d *dnsServer) Stop() {
|
|||||||
d.shutdownServer(srv, started, "stop")
|
d.shutdownServer(srv, started, "stop")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Query returns the address for the given name and query type. The second
|
func (d *dnsServer) Query(q uint16, data string) netip.Addr {
|
||||||
// return value reports whether the name is known at all (in either A or AAAA),
|
|
||||||
// which lets callers distinguish NODATA from NXDOMAIN.
|
|
||||||
func (d *dnsServer) Query(q uint16, data string) (netip.Addr, bool) {
|
|
||||||
data = strings.ToLower(data)
|
data = strings.ToLower(data)
|
||||||
d.RLock()
|
d.RLock()
|
||||||
defer d.RUnlock()
|
defer d.RUnlock()
|
||||||
addr4, haveV4 := d.dnsMap4[data]
|
|
||||||
addr6, haveV6 := d.dnsMap6[data]
|
|
||||||
nameExists := haveV4 || haveV6
|
|
||||||
switch q {
|
switch q {
|
||||||
case dns.TypeA:
|
case dns.TypeA:
|
||||||
if haveV4 {
|
if r, ok := d.dnsMap4[data]; ok {
|
||||||
return addr4, nameExists
|
return r
|
||||||
}
|
}
|
||||||
case dns.TypeAAAA:
|
case dns.TypeAAAA:
|
||||||
if haveV6 {
|
if r, ok := d.dnsMap6[data]; ok {
|
||||||
return addr6, nameExists
|
return r
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return netip.Addr{}, nameExists
|
return netip.Addr{}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *dnsServer) QueryCert(data string) string {
|
func (d *dnsServer) QueryCert(data string) string {
|
||||||
if len(data) < 2 {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
ip, err := netip.ParseAddr(data[:len(data)-1])
|
ip, err := netip.ParseAddr(data[:len(data)-1])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
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 ""
|
||||||
@@ -289,60 +257,12 @@ func (d *dnsServer) QueryCert(data string) string {
|
|||||||
return string(b)
|
return string(b)
|
||||||
}
|
}
|
||||||
|
|
||||||
// clearRecords drops all DNS records, including the self entry.
|
// clearRecords drops all DNS records.
|
||||||
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`
|
||||||
@@ -380,32 +300,17 @@ 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 cs.myVpnAddrsTable.Contains(b)
|
return d.myVpnAddrsTable.Contains(b)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
||||||
debugEnabled := d.l.Enabled(context.Background(), slog.LevelDebug)
|
|
||||||
// Per RFC 2308 §2.2, a name that exists but has no record of the requested
|
|
||||||
// type must be answered with NOERROR and an empty answer section (NODATA),
|
|
||||||
// not NXDOMAIN (RFC 2308 §2.1), which is reserved for names that do not
|
|
||||||
// exist at all.
|
|
||||||
anyNameExists := false
|
|
||||||
for _, q := range m.Question {
|
for _, q := range m.Question {
|
||||||
switch q.Qtype {
|
switch q.Qtype {
|
||||||
case dns.TypeA, dns.TypeAAAA:
|
case dns.TypeA, dns.TypeAAAA:
|
||||||
qType := dns.TypeToString[q.Qtype]
|
qType := dns.TypeToString[q.Qtype]
|
||||||
if debugEnabled {
|
d.l.Debugf("Query for %s %s", qType, q.Name)
|
||||||
d.l.Debug("DNS query", "type", qType, "name", q.Name)
|
ip := d.Query(q.Qtype, q.Name)
|
||||||
}
|
|
||||||
ip, nameExists := d.Query(q.Qtype, q.Name)
|
|
||||||
if nameExists {
|
|
||||||
anyNameExists = true
|
|
||||||
}
|
|
||||||
if ip.IsValid() {
|
if ip.IsValid() {
|
||||||
rr, err := dns.NewRR(fmt.Sprintf("%s %s %s", q.Name, qType, ip))
|
rr, err := dns.NewRR(fmt.Sprintf("%s %s %s", q.Name, qType, ip))
|
||||||
if err == nil {
|
if err == nil {
|
||||||
@@ -417,9 +322,7 @@ func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
|||||||
if !d.isSelfNebulaOrLocalhost(w.RemoteAddr().String()) {
|
if !d.isSelfNebulaOrLocalhost(w.RemoteAddr().String()) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if debugEnabled {
|
d.l.Debugf("Query for TXT %s", q.Name)
|
||||||
d.l.Debug("DNS query", "type", "TXT", "name", q.Name)
|
|
||||||
}
|
|
||||||
ip := d.QueryCert(q.Name)
|
ip := d.QueryCert(q.Name)
|
||||||
if ip != "" {
|
if ip != "" {
|
||||||
rr, err := dns.NewRR(fmt.Sprintf("%s TXT %s", q.Name, ip))
|
rr, err := dns.NewRR(fmt.Sprintf("%s TXT %s", q.Name, ip))
|
||||||
@@ -430,7 +333,7 @@ func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(m.Answer) == 0 && !anyNameExists {
|
if len(m.Answer) == 0 {
|
||||||
m.Rcode = dns.RcodeNameError
|
m.Rcode = dns.RcodeNameError
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+11
-351
@@ -2,37 +2,22 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"log/slog"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"strconv"
|
"strconv"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
)
|
)
|
||||||
|
|
||||||
type stubDNSWriter struct{}
|
|
||||||
|
|
||||||
func (stubDNSWriter) LocalAddr() net.Addr { return &net.UDPAddr{} }
|
|
||||||
func (stubDNSWriter) RemoteAddr() net.Addr {
|
|
||||||
return &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 5353}
|
|
||||||
}
|
|
||||||
func (stubDNSWriter) Write([]byte) (int, error) { return 0, nil }
|
|
||||||
func (stubDNSWriter) WriteMsg(*dns.Msg) error { return nil }
|
|
||||||
func (stubDNSWriter) Close() error { return nil }
|
|
||||||
func (stubDNSWriter) TsigStatus() error { return nil }
|
|
||||||
func (stubDNSWriter) TsigTimersOnly(bool) {}
|
|
||||||
func (stubDNSWriter) Hijack() {}
|
|
||||||
|
|
||||||
func TestParsequery(t *testing.T) {
|
func TestParsequery(t *testing.T) {
|
||||||
l := slog.New(slog.DiscardHandler)
|
l := logrus.New()
|
||||||
hostMap := &HostMap{}
|
hostMap := &HostMap{}
|
||||||
ds := &dnsServer{
|
ds := &dnsServer{
|
||||||
l: l,
|
l: l,
|
||||||
@@ -48,56 +33,18 @@ func TestParsequery(t *testing.T) {
|
|||||||
netip.MustParseAddr("fd01::25"),
|
netip.MustParseAddr("fd01::25"),
|
||||||
}
|
}
|
||||||
ds.Add("test.com.com", addrs)
|
ds.Add("test.com.com", addrs)
|
||||||
ds.Add("v4only.com.com", []netip.Addr{netip.MustParseAddr("1.2.3.6")})
|
|
||||||
ds.Add("v6only.com.com", []netip.Addr{netip.MustParseAddr("fd01::26")})
|
|
||||||
|
|
||||||
m := &dns.Msg{}
|
m := &dns.Msg{}
|
||||||
m.SetQuestion("test.com.com", dns.TypeA)
|
m.SetQuestion("test.com.com", dns.TypeA)
|
||||||
ds.parseQuery(m, nil)
|
ds.parseQuery(m, nil)
|
||||||
assert.NotNil(t, m.Answer)
|
assert.NotNil(t, m.Answer)
|
||||||
assert.Equal(t, "1.2.3.4", m.Answer[0].(*dns.A).A.String())
|
assert.Equal(t, "1.2.3.4", m.Answer[0].(*dns.A).A.String())
|
||||||
assert.Equal(t, dns.RcodeSuccess, m.Rcode)
|
|
||||||
|
|
||||||
m = &dns.Msg{}
|
m = &dns.Msg{}
|
||||||
m.SetQuestion("test.com.com", dns.TypeAAAA)
|
m.SetQuestion("test.com.com", dns.TypeAAAA)
|
||||||
ds.parseQuery(m, nil)
|
ds.parseQuery(m, nil)
|
||||||
assert.NotNil(t, m.Answer)
|
assert.NotNil(t, m.Answer)
|
||||||
assert.Equal(t, "fd01::24", m.Answer[0].(*dns.AAAA).AAAA.String())
|
assert.Equal(t, "fd01::24", m.Answer[0].(*dns.AAAA).AAAA.String())
|
||||||
assert.Equal(t, dns.RcodeSuccess, m.Rcode)
|
|
||||||
|
|
||||||
// A known name with no record of the requested type should return NODATA
|
|
||||||
// (NOERROR with empty answer), not NXDOMAIN.
|
|
||||||
m = &dns.Msg{}
|
|
||||||
m.SetQuestion("v4only.com.com", dns.TypeAAAA)
|
|
||||||
ds.parseQuery(m, nil)
|
|
||||||
assert.Empty(t, m.Answer)
|
|
||||||
assert.Equal(t, dns.RcodeSuccess, m.Rcode)
|
|
||||||
|
|
||||||
m = &dns.Msg{}
|
|
||||||
m.SetQuestion("v6only.com.com", dns.TypeA)
|
|
||||||
ds.parseQuery(m, nil)
|
|
||||||
assert.Empty(t, m.Answer)
|
|
||||||
assert.Equal(t, dns.RcodeSuccess, m.Rcode)
|
|
||||||
|
|
||||||
// An unknown name should still return NXDOMAIN.
|
|
||||||
m = &dns.Msg{}
|
|
||||||
m.SetQuestion("unknown.com.com", dns.TypeA)
|
|
||||||
ds.parseQuery(m, nil)
|
|
||||||
assert.Empty(t, m.Answer)
|
|
||||||
assert.Equal(t, dns.RcodeNameError, m.Rcode)
|
|
||||||
|
|
||||||
// short lookups should not fail
|
|
||||||
m = &dns.Msg{}
|
|
||||||
m.Question = []dns.Question{{Name: "", Qtype: dns.TypeTXT, Qclass: dns.ClassINET}}
|
|
||||||
ds.parseQuery(m, stubDNSWriter{})
|
|
||||||
assert.Empty(t, m.Answer)
|
|
||||||
assert.Equal(t, dns.RcodeNameError, m.Rcode)
|
|
||||||
|
|
||||||
m = &dns.Msg{}
|
|
||||||
m.Question = []dns.Question{{Name: ".", Qtype: dns.TypeTXT, Qclass: dns.ClassINET}}
|
|
||||||
ds.parseQuery(m, stubDNSWriter{})
|
|
||||||
assert.Empty(t, m.Answer)
|
|
||||||
assert.Equal(t, dns.RcodeNameError, m.Rcode)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func Test_getDnsServerAddr(t *testing.T) {
|
func Test_getDnsServerAddr(t *testing.T) {
|
||||||
@@ -139,9 +86,10 @@ func Test_getDnsServerAddr(t *testing.T) {
|
|||||||
|
|
||||||
func newTestDnsServer(t *testing.T) (*dnsServer, *config.C) {
|
func newTestDnsServer(t *testing.T) (*dnsServer, *config.C) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
sl := slog.New(slog.DiscardHandler)
|
l := logrus.New()
|
||||||
|
l.Out = io.Discard
|
||||||
ds := &dnsServer{
|
ds := &dnsServer{
|
||||||
l: sl,
|
l: l,
|
||||||
ctx: context.Background(),
|
ctx: context.Background(),
|
||||||
dnsMap4: make(map[string]netip.Addr),
|
dnsMap4: make(map[string]netip.Addr),
|
||||||
dnsMap6: make(map[string]netip.Addr),
|
dnsMap6: make(map[string]netip.Addr),
|
||||||
@@ -149,7 +97,7 @@ func newTestDnsServer(t *testing.T) (*dnsServer, *config.C) {
|
|||||||
}
|
}
|
||||||
ds.mux = dns.NewServeMux()
|
ds.mux = dns.NewServeMux()
|
||||||
ds.mux.HandleFunc(".", ds.handleDnsRequest)
|
ds.mux.HandleFunc(".", ds.handleDnsRequest)
|
||||||
return ds, config.NewC(nil)
|
return ds, config.NewC(l)
|
||||||
}
|
}
|
||||||
|
|
||||||
func setDnsConfig(c *config.C, host string, port string, amLighthouse, serveDns bool) {
|
func setDnsConfig(c *config.C, host string, port string, amLighthouse, serveDns bool) {
|
||||||
@@ -194,51 +142,14 @@ 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", port, true, true)
|
setDnsConfig(c, "127.0.0.1", "0", 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) {
|
||||||
@@ -316,92 +227,6 @@ 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)
|
||||||
@@ -464,168 +289,3 @@ 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,16 +1,6 @@
|
|||||||
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
|
||||||
|
|
||||||
|
|||||||
@@ -1,85 +0,0 @@
|
|||||||
//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
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -28,7 +28,6 @@ func makeHandshakePacket(from, to netip.AddrPort, subtype header.MessageSubType,
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify the responder correctly handles receiving the same msg1 multiple times
|
// Verify the responder correctly handles receiving the same msg1 multiple times
|
||||||
// (retransmission). The duplicate goes through CheckAndComplete -> ErrAlreadySeen
|
// (retransmission). The duplicate goes through CheckAndComplete -> ErrAlreadySeen
|
||||||
// and the cached response is resent.
|
// and the cached response is resent.
|
||||||
@@ -47,7 +46,7 @@ func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me to them")
|
t.Log("Trigger handshake from me to them")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
|
|
||||||
t.Log("Grab my msg1")
|
t.Log("Grab my msg1")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -79,7 +78,6 @@ func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify that a truncated handshake packet is ignored and the real
|
// Verify that a truncated handshake packet is ignored and the real
|
||||||
// packet can still complete the handshake.
|
// packet can still complete the handshake.
|
||||||
|
|
||||||
@@ -97,7 +95,7 @@ func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake")
|
t.Log("Trigger handshake")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
|
|
||||||
t.Log("Get msg1 and deliver to responder")
|
t.Log("Get msg1 and deliver to responder")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -128,7 +126,6 @@ func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// A msg2 arriving with no matching pending index should be silently dropped
|
// A msg2 arriving with no matching pending index should be silently dropped
|
||||||
// with no response sent and no state changes.
|
// with no response sent and no state changes.
|
||||||
|
|
||||||
@@ -146,7 +143,7 @@ func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Complete a normal handshake")
|
t.Log("Complete a normal handshake")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
@@ -171,7 +168,6 @@ func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// A handshake packet with an unexpected message counter should be silently
|
// A handshake packet with an unexpected message counter should be silently
|
||||||
// dropped with no side effects and no UDP response.
|
// dropped with no side effects and no UDP response.
|
||||||
|
|
||||||
@@ -203,7 +199,6 @@ func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeUnknownSubtype(t *testing.T) {
|
func TestHandshakeUnknownSubtype(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// A handshake packet with an unknown subtype should be silently dropped.
|
// A handshake packet with an unknown subtype should be silently dropped.
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -229,7 +224,6 @@ func TestHandshakeUnknownSubtype(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeLateResponse(t *testing.T) {
|
func TestHandshakeLateResponse(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// After a handshake times out, a late response should be silently ignored
|
// After a handshake times out, a late response should be silently ignored
|
||||||
// with no new tunnels created.
|
// with no new tunnels created.
|
||||||
|
|
||||||
@@ -248,7 +242,7 @@ func TestHandshakeLateResponse(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me")
|
t.Log("Trigger handshake from me")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
|
|
||||||
t.Log("Grab msg1 but don't deliver")
|
t.Log("Grab msg1 but don't deliver")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -279,7 +273,6 @@ func TestHandshakeLateResponse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify that a node rejects a handshake containing its own VPN IP in the
|
// Verify that a node rejects a handshake containing its own VPN IP in the
|
||||||
// peer cert. We do this by sending the initiator's own msg1 back to itself.
|
// peer cert. We do this by sending the initiator's own msg1 back to itself.
|
||||||
|
|
||||||
@@ -292,7 +285,7 @@ func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
|||||||
myControl.Start()
|
myControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me")
|
t.Log("Trigger handshake from me")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(netip.MustParseAddr("10.128.0.2"), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(netip.MustParseAddr("10.128.0.2"), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
|
|
||||||
t.Log("Drain any handshake retransmits before injecting")
|
t.Log("Drain any handshake retransmits before injecting")
|
||||||
@@ -328,7 +321,6 @@ func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// MessageCounter=0 is not a valid handshake message and should be dropped.
|
// MessageCounter=0 is not a valid handshake message and should be dropped.
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -349,7 +341,6 @@ func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeRemoteAllowList(t *testing.T) {
|
func TestHandshakeRemoteAllowList(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify that a handshake from a blocked underlay IP is dropped with no
|
// Verify that a handshake from a blocked underlay IP is dropped with no
|
||||||
// response and no state changes. Then verify the same packet from an
|
// response and no state changes. Then verify the same packet from an
|
||||||
// allowed IP succeeds.
|
// allowed IP succeeds.
|
||||||
@@ -375,7 +366,7 @@ func TestHandshakeRemoteAllowList(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake from them")
|
t.Log("Trigger handshake from them")
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
msg1 := theirControl.GetFromUDP(true)
|
msg1 := theirControl.GetFromUDP(true)
|
||||||
|
|
||||||
t.Log("Rewrite the source to a blocked IP and inject")
|
t.Log("Rewrite the source to a blocked IP and inject")
|
||||||
@@ -408,7 +399,6 @@ func TestHandshakeRemoteAllowList(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// When a duplicate msg1 arrives via ErrAlreadySeen, verify the tunnel
|
// When a duplicate msg1 arrives via ErrAlreadySeen, verify the tunnel
|
||||||
// remains functional and hostmap index count is stable.
|
// remains functional and hostmap index count is stable.
|
||||||
|
|
||||||
@@ -426,7 +416,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Complete a normal handshake via the router")
|
t.Log("Complete a normal handshake via the router")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
@@ -437,7 +427,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|||||||
originalRemote := hi.CurrentRemote
|
originalRemote := hi.CurrentRemote
|
||||||
|
|
||||||
t.Log("Re-trigger traffic to cause a new handshake attempt (ErrAlreadySeen)")
|
t.Log("Re-trigger traffic to cause a new handshake attempt (ErrAlreadySeen)")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("roam")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("roam"))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
|
|
||||||
t.Log("Verify tunnel still works")
|
t.Log("Verify tunnel still works")
|
||||||
@@ -455,7 +445,6 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify that when the wrong host responds, the cached packets are
|
// Verify that when the wrong host responds, the cached packets are
|
||||||
// transferred to the new handshake, the evil tunnel is closed, evil's
|
// transferred to the new handshake, the evil tunnel is closed, evil's
|
||||||
// address is blocked, and the correct tunnel is eventually established.
|
// address is blocked, and the correct tunnel is eventually established.
|
||||||
@@ -475,8 +464,8 @@ func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Send multiple packets to them (cached during handshake)")
|
t.Log("Send multiple packets to them (cached during handshake)")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet1")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet1"))
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet2")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet2"))
|
||||||
|
|
||||||
t.Log("Route until evil tunnel is closed")
|
t.Log("Route until evil tunnel is closed")
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
@@ -519,7 +508,6 @@ func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeRelayComplete(t *testing.T) {
|
func TestHandshakeRelayComplete(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify that a relay handshake completes correctly and relay state is
|
// Verify that a relay handshake completes correctly and relay state is
|
||||||
// properly maintained on all three nodes.
|
// properly maintained on all three nodes.
|
||||||
|
|
||||||
@@ -540,7 +528,7 @@ func TestHandshakeRelayComplete(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake via relay")
|
t.Log("Trigger handshake via relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi via relay")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi via relay"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi via relay"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi via relay"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
@@ -568,7 +556,7 @@ func TestHandshakeRelayComplete(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// NOTE: Relay V1 cert + IPv6 rejection is not tested here because
|
// NOTE: Relay V1 cert + IPv6 rejection is not tested here because
|
||||||
// BuildTunUDPPacket from a V4 node to a V6 address panics in the test
|
// InjectTunUDPPacket from a V4 node to a V6 address panics in the test
|
||||||
// framework. The check is in handshake_manager.go handleOutbound relay
|
// framework. The check is in handshake_manager.go handleOutbound relay
|
||||||
// logic (lines ~304-313): if the relay host has a V1 cert and either
|
// logic (lines ~304-313): if the relay host has a V1 cert and either
|
||||||
// address is IPv6, the relay is skipped.
|
// address is IPv6, the relay is skipped.
|
||||||
|
|||||||
+77
-321
@@ -11,12 +11,12 @@ import (
|
|||||||
|
|
||||||
"github.com/google/gopacket"
|
"github.com/google/gopacket"
|
||||||
"github.com/google/gopacket/layers"
|
"github.com/google/gopacket/layers"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/cert_test"
|
"github.com/slackhq/nebula/cert_test"
|
||||||
"github.com/slackhq/nebula/e2e/router"
|
"github.com/slackhq/nebula/e2e/router"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -40,22 +40,11 @@ func BenchmarkHotPath(b *testing.B) {
|
|||||||
r.CancelFlowLogs()
|
r.CancelFlowLogs()
|
||||||
|
|
||||||
assertTunnel(b, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(b, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
// Pre-build the IP packet bytes once so the bench measures the data plane,
|
|
||||||
// not gopacket SerializeLayers overhead.
|
|
||||||
prebuilt := BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
|
||||||
|
|
||||||
// EnableFanIn switches the router to a 0-alloc routing path. Required
|
|
||||||
// for hot-path benchmarks; would conflict with GetFromUDP-using tests.
|
|
||||||
r.EnableFanIn()
|
|
||||||
|
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
|
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
myControl.InjectTunPacket(prebuilt)
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
// Release the TUN-side bytes back to the harness freelist; the bench
|
_ = r.RouteForAllUntilTxTun(theirControl)
|
||||||
// just confirms a packet arrived, the contents aren't inspected.
|
|
||||||
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
@@ -83,15 +72,11 @@ func BenchmarkHotPathRelay(b *testing.B) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
assertTunnel(b, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(b, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
|
|
||||||
prebuilt := BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
|
||||||
r.EnableFanIn()
|
|
||||||
|
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
|
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
myControl.InjectTunPacket(prebuilt)
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
_ = r.RouteForAllUntilTxTun(theirControl)
|
||||||
}
|
}
|
||||||
|
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
@@ -100,7 +85,6 @@ func BenchmarkHotPathRelay(b *testing.B) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGoodHandshake(t *testing.T) {
|
func TestGoodHandshake(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -113,7 +97,7 @@ func TestGoodHandshake(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||||
@@ -151,7 +135,6 @@ func TestGoodHandshake(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGoodHandshakeNoOverlap(t *testing.T) {
|
func TestGoodHandshakeNoOverlap(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "2001::69/24", nil) //look ma, cross-stack!
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "2001::69/24", nil) //look ma, cross-stack!
|
||||||
@@ -187,7 +170,6 @@ func TestGoodHandshakeNoOverlap(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestWrongResponderHandshake(t *testing.T) {
|
func TestWrongResponderHandshake(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.100/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.100/24", nil)
|
||||||
@@ -207,7 +189,7 @@ func TestWrongResponderHandshake(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
@@ -264,7 +246,6 @@ func TestWrongResponderHandshake(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.99/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.99/24", nil)
|
||||||
@@ -289,7 +270,7 @@ func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
@@ -347,7 +328,6 @@ func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestStage1Race(t *testing.T) {
|
func TestStage1Race(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// This tests ensures that two hosts handshaking with each other at the same time will allow traffic to flow
|
// This tests ensures that two hosts handshaking with each other at the same time will allow traffic to flow
|
||||||
// But will eventually collapse down to a single tunnel
|
// But will eventually collapse down to a single tunnel
|
||||||
|
|
||||||
@@ -368,8 +348,8 @@ func TestStage1Race(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake to start on both me and them")
|
t.Log("Trigger a handshake to start on both me and them")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
t.Log("Get both stage 1 handshake packets")
|
t.Log("Get both stage 1 handshake packets")
|
||||||
myHsForThem := myControl.GetFromUDP(true)
|
myHsForThem := myControl.GetFromUDP(true)
|
||||||
@@ -405,7 +385,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 myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 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)
|
||||||
@@ -428,7 +408,6 @@ func TestStage1Race(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUncleanShutdownRaceLoser(t *testing.T) {
|
func TestUncleanShutdownRaceLoser(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -446,20 +425,18 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
r.Log("Trigger a handshake from me to them")
|
r.Log("Trigger a handshake from me to them")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
|
||||||
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.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again"))
|
||||||
p = r.RouteForAllUntilTxTun(theirControl)
|
p = r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me again"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me again"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
|
||||||
@@ -467,10 +444,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 := theirControl.GetHostmapIndexCount()
|
start := len(theirControl.GetHostmap().Indexes)
|
||||||
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 theirControl.GetHostmapIndexCount() < start {
|
if len(theirControl.GetHostmap().Indexes) < start {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
@@ -480,7 +457,6 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUncleanShutdownRaceWinner(t *testing.T) {
|
func TestUncleanShutdownRaceWinner(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -498,7 +474,7 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
r.Log("Trigger a handshake from me to them")
|
r.Log("Trigger a handshake from me to them")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
@@ -506,13 +482,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.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again"))
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
assertUdpPacket(t, []byte("Hi from them again"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from them again"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
||||||
r.RenderHostmaps("Derp hostmaps", myControl, theirControl)
|
r.RenderHostmaps("Derp hostmaps", myControl, theirControl)
|
||||||
@@ -521,10 +495,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 := myControl.GetHostmapIndexCount()
|
start := len(myControl.GetHostmap().Indexes)
|
||||||
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 myControl.GetHostmapIndexCount() < start {
|
if len(myControl.GetHostmap().Indexes) < start {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
@@ -534,7 +508,6 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRelays(t *testing.T) {
|
func TestRelays(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -555,7 +528,7 @@ func TestRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -564,7 +537,6 @@ func TestRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRelaysDontCareAboutIps(t *testing.T) {
|
func TestRelaysDontCareAboutIps(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "2001::9999/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "2001::9999/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -585,7 +557,7 @@ func TestRelaysDontCareAboutIps(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -594,7 +566,6 @@ func TestRelaysDontCareAboutIps(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestReestablishRelays(t *testing.T) {
|
func TestReestablishRelays(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -615,14 +586,14 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
|
||||||
t.Log("Ensure packet traversal from them to me via the relay")
|
t.Log("Ensure packet traversal from them to me via the relay")
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -632,12 +603,12 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
r.Log("Close the tunnel")
|
r.Log("Close the tunnel")
|
||||||
relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
|
relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
|
||||||
|
|
||||||
start := myControl.GetHostmapIndexCount()
|
start := len(myControl.GetHostmap().Indexes)
|
||||||
curIndexes := myControl.GetHostmapIndexCount()
|
curIndexes := len(myControl.GetHostmap().Indexes)
|
||||||
for curIndexes >= start {
|
for curIndexes >= start {
|
||||||
curIndexes = myControl.GetHostmapIndexCount()
|
curIndexes = len(myControl.GetHostmap().Indexes)
|
||||||
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail"))
|
||||||
|
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
return router.RouteAndExit
|
return router.RouteAndExit
|
||||||
@@ -654,7 +625,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(theirControl)
|
p = r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -689,7 +660,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
t.Log("Assert the tunnel works the other way, too")
|
t.Log("Assert the tunnel works the other way, too")
|
||||||
for {
|
for {
|
||||||
t.Log("RouteForAllUntilTxTun")
|
t.Log("RouteForAllUntilTxTun")
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -725,72 +696,7 @@ 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()
|
|
||||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
@@ -823,8 +729,8 @@ func TestStage1RaceRelays(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
||||||
|
|
||||||
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
r.Log("Wait for a packet from them to me")
|
r.Log("Wait for a packet from them to me")
|
||||||
p := r.RouteForAllUntilTxTun(myControl)
|
p := r.RouteForAllUntilTxTun(myControl)
|
||||||
@@ -838,12 +744,12 @@ func TestStage1RaceRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestStage1RaceRelays2(t *testing.T) {
|
func TestStage1RaceRelays2(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them ", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them ", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
|
||||||
|
l := NewTestLogger()
|
||||||
|
|
||||||
// Teach my how to get to the relay and that their can be reached via the relay
|
// Teach my how to get to the relay and that their can be reached via the relay
|
||||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||||
@@ -865,41 +771,49 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
r.Log("Get a tunnel between me and relay")
|
r.Log("Get a tunnel between me and relay")
|
||||||
|
l.Info("Get a tunnel between me and relay")
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), myControl, relayControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), myControl, relayControl, r)
|
||||||
|
|
||||||
r.Log("Get a tunnel between them and relay")
|
r.Log("Get a tunnel between them and relay")
|
||||||
|
l.Info("Get a tunnel between them and relay")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
||||||
|
|
||||||
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
l.Info("Trigger a handshake from both them and me via relay to them and me")
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
//r.RouteUntilAfterMsgType(myControl, header.Control, header.MessageNone)
|
//r.RouteUntilAfterMsgType(myControl, header.Control, header.MessageNone)
|
||||||
//r.RouteUntilAfterMsgType(theirControl, header.Control, header.MessageNone)
|
//r.RouteUntilAfterMsgType(theirControl, header.Control, header.MessageNone)
|
||||||
|
|
||||||
r.Log("Wait for a packet from them to me; myControl")
|
r.Log("Wait for a packet from them to me")
|
||||||
|
l.Info("Wait for a packet from them to me; myControl")
|
||||||
r.RouteForAllUntilTxTun(myControl)
|
r.RouteForAllUntilTxTun(myControl)
|
||||||
r.Log("Wait for a packet from them to me; theirControl")
|
l.Info("Wait for a packet from them to me; theirControl")
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
|
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
|
l.Info("Assert the tunnel works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
|
|
||||||
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",
|
l.Info("Wait until we remove extra tunnels")
|
||||||
myControl.GetHostmapIndexCount(),
|
l.WithFields(
|
||||||
theirControl.GetHostmapIndexCount(),
|
logrus.Fields{
|
||||||
relayControl.GetHostmapIndexCount(),
|
"myControl": len(myControl.GetHostmap().Indexes),
|
||||||
)
|
"theirControl": len(theirControl.GetHostmap().Indexes),
|
||||||
hostInfos := myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount()
|
"relayControl": len(relayControl.GetHostmap().Indexes),
|
||||||
|
}).Info("Waiting for hostinfos to be removed...")
|
||||||
|
hostInfos := len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
|
||||||
retries := 60
|
retries := 60
|
||||||
for hostInfos > 6 && retries > 0 {
|
for hostInfos > 6 && retries > 0 {
|
||||||
hostInfos = myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount()
|
hostInfos = len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
|
||||||
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
l.WithFields(
|
||||||
myControl.GetHostmapIndexCount(),
|
logrus.Fields{
|
||||||
theirControl.GetHostmapIndexCount(),
|
"myControl": len(myControl.GetHostmap().Indexes),
|
||||||
relayControl.GetHostmapIndexCount(),
|
"theirControl": len(theirControl.GetHostmap().Indexes),
|
||||||
)
|
"relayControl": len(relayControl.GetHostmap().Indexes),
|
||||||
|
}).Info("Waiting for hostinfos to be removed...")
|
||||||
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)
|
||||||
@@ -907,6 +821,7 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
|
l.Info("Assert the tunnel works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
|
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
@@ -915,7 +830,6 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingRelays(t *testing.T) {
|
func TestRehandshakingRelays(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, relayConfig := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, relayConfig := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -936,7 +850,7 @@ func TestRehandshakingRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -992,24 +906,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 myControl.GetHostmapIndexCount() != 2 {
|
for len(myControl.GetHostmap().Indexes) != 2 {
|
||||||
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount())
|
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
|
||||||
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 theirControl.GetHostmapIndexCount() != 2 {
|
for len(theirControl.GetHostmap().Indexes) != 2 {
|
||||||
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount())
|
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
|
||||||
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 relayControl.GetHostmapIndexCount() != 2 {
|
for len(relayControl.GetHostmap().Indexes) != 2 {
|
||||||
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount())
|
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
|
||||||
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")
|
||||||
@@ -1019,7 +933,6 @@ func TestRehandshakingRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingRelaysPrimary(t *testing.T) {
|
func TestRehandshakingRelaysPrimary(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// This test is the same as TestRehandshakingRelays but one of the terminal types is a primary swap winner
|
// This test is the same as TestRehandshakingRelays but one of the terminal types is a primary swap winner
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.128/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.128/24", m{"relay": m{"use_relays": true}})
|
||||||
@@ -1041,7 +954,7 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -1097,24 +1010,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 myControl.GetHostmapIndexCount() != 2 {
|
for len(myControl.GetHostmap().Indexes) != 2 {
|
||||||
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount())
|
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
|
||||||
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 theirControl.GetHostmapIndexCount() != 2 {
|
for len(theirControl.GetHostmap().Indexes) != 2 {
|
||||||
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount())
|
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
|
||||||
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 relayControl.GetHostmapIndexCount() != 2 {
|
for len(relayControl.GetHostmap().Indexes) != 2 {
|
||||||
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount())
|
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
|
||||||
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")
|
||||||
@@ -1124,7 +1037,6 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshaking(t *testing.T) {
|
func TestRehandshaking(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, myConfig := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.2/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, myConfig := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.2/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, theirConfig := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.1/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, theirConfig := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.1/24", nil)
|
||||||
@@ -1191,7 +1103,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 myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 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)
|
||||||
@@ -1220,7 +1132,6 @@ func TestRehandshaking(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingLoser(t *testing.T) {
|
func TestRehandshakingLoser(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// The purpose of this test is that the race loser renews their certificate and rehandshakes. The final tunnel
|
// The purpose of this test is that the race loser renews their certificate and rehandshakes. The final tunnel
|
||||||
// Should be the one with the new certificate
|
// Should be the one with the new certificate
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -1291,7 +1202,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 myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 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)
|
||||||
@@ -1319,7 +1230,6 @@ func TestRehandshakingLoser(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRaceRegression(t *testing.T) {
|
func TestRaceRegression(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// This test forces stage 1, stage 2, stage 1 to be received by me from them
|
// This test forces stage 1, stage 2, stage 1 to be received by me from them
|
||||||
// We had a bug where we were not finding the duplicate handshake and responding to the final stage 1 which
|
// We had a bug where we were not finding the duplicate handshake and responding to the final stage 1 which
|
||||||
// caused a cross-linked hostinfo
|
// caused a cross-linked hostinfo
|
||||||
@@ -1343,8 +1253,8 @@ func TestRaceRegression(t *testing.T) {
|
|||||||
//them rx stage:2 initiatorIndex=120607833 responderIndex=4209862089
|
//them rx stage:2 initiatorIndex=120607833 responderIndex=4209862089
|
||||||
|
|
||||||
t.Log("Start both handshakes")
|
t.Log("Start both handshakes")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
t.Log("Get both stage 1")
|
t.Log("Get both stage 1")
|
||||||
myStage1ForThem := myControl.GetFromUDP(true)
|
myStage1ForThem := myControl.GetFromUDP(true)
|
||||||
@@ -1380,7 +1290,6 @@ func TestRaceRegression(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestV2NonPrimaryWithLighthouse(t *testing.T) {
|
func TestV2NonPrimaryWithLighthouse(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "10.128.0.1/24, ff::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "10.128.0.1/24, ff::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
||||||
|
|
||||||
@@ -1421,7 +1330,6 @@ func TestV2NonPrimaryWithLighthouse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestV2NonPrimaryWithOffNetLighthouse(t *testing.T) {
|
func TestV2NonPrimaryWithOffNetLighthouse(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "2001::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "2001::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
||||||
|
|
||||||
@@ -1461,84 +1369,7 @@ func TestV2NonPrimaryWithOffNetLighthouse(t *testing.T) {
|
|||||||
theirControl.Stop()
|
theirControl.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestLighthouseUpdateOnReload(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
|
|
||||||
// Create the lighthouse
|
|
||||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{"lighthouse": m{"am_lighthouse": true}})
|
|
||||||
|
|
||||||
// Create a client with NO lighthouse configured and a long update interval.
|
|
||||||
// The initial SendUpdate at startup will be a no-op since no lighthouses are known.
|
|
||||||
myControl, myVpnIpNet, _, myConfig := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
|
||||||
"lighthouse": m{
|
|
||||||
"interval": 600,
|
|
||||||
"local_allow_list": m{
|
|
||||||
"10.0.0.0/24": true,
|
|
||||||
"::/0": false,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
r := router.NewR(t, lhControl, myControl)
|
|
||||||
defer r.RenderFlow()
|
|
||||||
|
|
||||||
lhControl.Start()
|
|
||||||
myControl.Start()
|
|
||||||
|
|
||||||
// Drain any startup packets (there should be none meaningful)
|
|
||||||
r.FlushAll()
|
|
||||||
|
|
||||||
// Verify lighthouse has no knowledge of the client
|
|
||||||
assert.Nil(t, lhControl.QueryLighthouse(myVpnIpNet[0].Addr()))
|
|
||||||
|
|
||||||
// Build a new config that adds the lighthouse
|
|
||||||
newSettings := make(m)
|
|
||||||
for k, v := range myConfig.Settings {
|
|
||||||
newSettings[k] = v
|
|
||||||
}
|
|
||||||
newSettings["static_host_map"] = m{
|
|
||||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
|
||||||
}
|
|
||||||
newSettings["lighthouse"] = m{
|
|
||||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
|
||||||
"interval": 600,
|
|
||||||
"local_allow_list": m{
|
|
||||||
"10.0.0.0/24": true,
|
|
||||||
"::/0": false,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
newCfg, err := yaml.Marshal(newSettings)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Reload the config. The lighthouse.hosts change triggers TriggerUpdate,
|
|
||||||
// which wakes the update worker. It calls SendUpdate, initiating a
|
|
||||||
// handshake to the new lighthouse and caching the HostUpdateNotification.
|
|
||||||
require.NoError(t, myConfig.ReloadConfigString(string(newCfg)))
|
|
||||||
|
|
||||||
// Route until the lighthouse receives the HostUpdateNotification.
|
|
||||||
// This covers: handshake stage 1, stage 2, then the cached update.
|
|
||||||
done := make(chan struct{})
|
|
||||||
go func() {
|
|
||||||
r.RouteForAllUntilAfterMsgTypeTo(lhControl, header.LightHouse, 0)
|
|
||||||
close(done)
|
|
||||||
}()
|
|
||||||
select {
|
|
||||||
case <-done:
|
|
||||||
case <-time.After(5 * time.Second):
|
|
||||||
t.Fatal("timed out waiting for lighthouse update after config reload")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify lighthouse now has the client's addresses
|
|
||||||
assert.NotNil(t, lhControl.QueryLighthouse(myVpnIpNet[0].Addr()))
|
|
||||||
|
|
||||||
r.RenderHostmaps("Final hostmaps", lhControl, myControl)
|
|
||||||
lhControl.Stop()
|
|
||||||
myControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
unsafePrefix := "192.168.6.0/24"
|
unsafePrefix := "192.168.6.0/24"
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServerWithUdpAndUnsafeNetworks(cert.Version2, ca, caKey, "spooky", "10.128.0.2/24", netip.MustParseAddrPort("10.64.0.2:4242"), unsafePrefix, nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServerWithUdpAndUnsafeNetworks(cert.Version2, ca, caKey, "spooky", "10.128.0.2/24", netip.MustParseAddrPort("10.64.0.2:4242"), unsafePrefix, nil)
|
||||||
@@ -1560,7 +1391,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||||
@@ -1588,7 +1419,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
|||||||
assertUdpPacket(t, []byte("Hi from me"), myCachedPacket, myVpnIpNet[0].Addr(), spookyDest, 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), myCachedPacket, myVpnIpNet[0].Addr(), spookyDest, 80, 80)
|
||||||
|
|
||||||
//reply
|
//reply
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman"))
|
||||||
//wait for reply
|
//wait for reply
|
||||||
theirControl.WaitForType(1, 0, myControl)
|
theirControl.WaitForType(1, 0, myControl)
|
||||||
theirCachedPacket := myControl.GetFromTun(true)
|
theirCachedPacket := myControl.GetFromTun(true)
|
||||||
@@ -1603,78 +1434,3 @@ 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()
|
|
||||||
}
|
|
||||||
|
|||||||
+20
-83
@@ -4,6 +4,7 @@
|
|||||||
package e2e
|
package e2e
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -11,18 +12,15 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"log/slog"
|
|
||||||
|
|
||||||
"dario.cat/mergo"
|
"dario.cat/mergo"
|
||||||
"github.com/google/gopacket"
|
"github.com/google/gopacket"
|
||||||
"github.com/google/gopacket/layers"
|
"github.com/google/gopacket/layers"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/cert_test"
|
"github.com/slackhq/nebula/cert_test"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/e2e/router"
|
"github.com/slackhq/nebula/e2e/router"
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"go.yaml.in/yaml/v3"
|
"go.yaml.in/yaml/v3"
|
||||||
@@ -134,7 +132,8 @@ func newSimpleServerWithUdpAndUnsafeNetworks(v cert.Version, caCrt cert.Certific
|
|||||||
"port": udpAddr.Port(),
|
"port": udpAddr.Port(),
|
||||||
},
|
},
|
||||||
"logging": m{
|
"logging": m{
|
||||||
"level": testLogLevelName(),
|
"timestamp_format": fmt.Sprintf("%v 15:04:05.000000", name),
|
||||||
|
"level": l.Level.String(),
|
||||||
},
|
},
|
||||||
"timers": m{
|
"timers": m{
|
||||||
"pending_deletion_interval": 2,
|
"pending_deletion_interval": 2,
|
||||||
@@ -235,7 +234,8 @@ func newServer(caCrt []cert.Certificate, certs []cert.Certificate, key []byte, o
|
|||||||
"port": udpAddr.Port(),
|
"port": udpAddr.Port(),
|
||||||
},
|
},
|
||||||
"logging": m{
|
"logging": m{
|
||||||
"level": testLogLevelName(),
|
"timestamp_format": fmt.Sprintf("%v 15:04:05.000000", certs[0].Name()),
|
||||||
|
"level": l.Level.String(),
|
||||||
},
|
},
|
||||||
"timers": m{
|
"timers": m{
|
||||||
"pending_deletion_interval": 2,
|
"pending_deletion_interval": 2,
|
||||||
@@ -294,12 +294,12 @@ func deadline(t *testing.T, seconds time.Duration) doneCb {
|
|||||||
|
|
||||||
func assertTunnel(t testing.TB, vpnIpA, vpnIpB netip.Addr, controlA, controlB *nebula.Control, r *router.R) {
|
func assertTunnel(t testing.TB, vpnIpA, vpnIpB netip.Addr, controlA, controlB *nebula.Control, r *router.R) {
|
||||||
// Send a packet from them to me
|
// Send a packet from them to me
|
||||||
controlB.InjectTunPacket(BuildTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B")))
|
controlB.InjectTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B"))
|
||||||
bPacket := r.RouteForAllUntilTxTun(controlA)
|
bPacket := r.RouteForAllUntilTxTun(controlA)
|
||||||
assertUdpPacket(t, []byte("Hi from B"), bPacket, vpnIpB, vpnIpA, 90, 80)
|
assertUdpPacket(t, []byte("Hi from B"), bPacket, vpnIpB, vpnIpA, 90, 80)
|
||||||
|
|
||||||
// And once more from me to them
|
// And once more from me to them
|
||||||
controlA.InjectTunPacket(BuildTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A")))
|
controlA.InjectTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A"))
|
||||||
aPacket := r.RouteForAllUntilTxTun(controlB)
|
aPacket := r.RouteForAllUntilTxTun(controlB)
|
||||||
assertUdpPacket(t, []byte("Hello from A"), aPacket, vpnIpA, vpnIpB, 90, 80)
|
assertUdpPacket(t, []byte("Hello from A"), aPacket, vpnIpA, vpnIpB, 90, 80)
|
||||||
}
|
}
|
||||||
@@ -379,87 +379,24 @@ func getAddrs(ns []netip.Prefix) []netip.Addr {
|
|||||||
return a
|
return a
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewTestLogger() *slog.Logger {
|
func NewTestLogger() *logrus.Logger {
|
||||||
|
l := logrus.New()
|
||||||
|
|
||||||
v := os.Getenv("TEST_LOGS")
|
v := os.Getenv("TEST_LOGS")
|
||||||
if v == "" {
|
if v == "" {
|
||||||
return slog.New(slog.NewTextHandler(io.Discard, nil))
|
l.SetOutput(io.Discard)
|
||||||
|
l.SetLevel(logrus.PanicLevel)
|
||||||
|
return l
|
||||||
}
|
}
|
||||||
|
|
||||||
level := slog.LevelInfo
|
|
||||||
switch v {
|
switch v {
|
||||||
case "2":
|
case "2":
|
||||||
level = slog.LevelDebug
|
l.SetLevel(logrus.DebugLevel)
|
||||||
case "3":
|
case "3":
|
||||||
level = logging.LevelTrace
|
l.SetLevel(logrus.TraceLevel)
|
||||||
|
default:
|
||||||
|
l.SetLevel(logrus.InfoLevel)
|
||||||
}
|
}
|
||||||
return slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: level}))
|
|
||||||
}
|
return l
|
||||||
|
|
||||||
// testLogLevelName returns the level name string accepted by logging.ApplyConfig
|
|
||||||
// for the current TEST_LOGS setting. Kept in sync with NewTestLogger.
|
|
||||||
func testLogLevelName() string {
|
|
||||||
switch os.Getenv("TEST_LOGS") {
|
|
||||||
case "2":
|
|
||||||
return "debug"
|
|
||||||
case "3":
|
|
||||||
return "trace"
|
|
||||||
case "":
|
|
||||||
return "info"
|
|
||||||
}
|
|
||||||
return "info"
|
|
||||||
}
|
|
||||||
|
|
||||||
// BuildTunUDPPacket assembles an IP+UDP packet suitable for Control.InjectTunPacket.
|
|
||||||
// Using UDP here because it's a simpler protocol.
|
|
||||||
func BuildTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr netip.Addr, fromPort uint16, data []byte) []byte {
|
|
||||||
serialize := make([]gopacket.SerializableLayer, 0)
|
|
||||||
var netLayer gopacket.NetworkLayer
|
|
||||||
if toAddr.Is6() {
|
|
||||||
if !fromAddr.Is6() {
|
|
||||||
panic("Cant send ipv6 to ipv4")
|
|
||||||
}
|
|
||||||
ip := &layers.IPv6{
|
|
||||||
Version: 6,
|
|
||||||
NextHeader: layers.IPProtocolUDP,
|
|
||||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
|
||||||
DstIP: toAddr.Unmap().AsSlice(),
|
|
||||||
}
|
|
||||||
serialize = append(serialize, ip)
|
|
||||||
netLayer = ip
|
|
||||||
} else {
|
|
||||||
if !fromAddr.Is4() {
|
|
||||||
panic("Cant send ipv4 to ipv6")
|
|
||||||
}
|
|
||||||
|
|
||||||
ip := &layers.IPv4{
|
|
||||||
Version: 4,
|
|
||||||
TTL: 64,
|
|
||||||
Protocol: layers.IPProtocolUDP,
|
|
||||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
|
||||||
DstIP: toAddr.Unmap().AsSlice(),
|
|
||||||
}
|
|
||||||
serialize = append(serialize, ip)
|
|
||||||
netLayer = ip
|
|
||||||
}
|
|
||||||
|
|
||||||
udp := layers.UDP{
|
|
||||||
SrcPort: layers.UDPPort(fromPort),
|
|
||||||
DstPort: layers.UDPPort(toPort),
|
|
||||||
}
|
|
||||||
if err := udp.SetNetworkLayerForChecksum(netLayer); err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
buffer := gopacket.NewSerializeBuffer()
|
|
||||||
opt := gopacket.SerializeOptions{
|
|
||||||
ComputeChecksums: true,
|
|
||||||
FixLengths: true,
|
|
||||||
}
|
|
||||||
|
|
||||||
serialize = append(serialize, &udp, gopacket.Payload(data))
|
|
||||||
if err := gopacket.SerializeLayers(buffer, opt, serialize...); err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return buffer.Bytes()
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,47 +0,0 @@
|
|||||||
//go:build e2e_testing
|
|
||||||
// +build e2e_testing
|
|
||||||
|
|
||||||
package e2e
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/cert_test"
|
|
||||||
"github.com/slackhq/nebula/e2e/router"
|
|
||||||
"go.uber.org/goleak"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestNoGoroutineLeaks brings up two nebula instances, completes a tunnel,
|
|
||||||
// stops both, and asserts no goroutines leak past the shutdown. goleak's
|
|
||||||
// retry mechanism gives the wg.Wait()-driven goroutines a moment to drain
|
|
||||||
// before failing the assertion.
|
|
||||||
//
|
|
||||||
// Intentionally NOT t.Parallel()'d: concurrent tests would have their own
|
|
||||||
// goroutines running and trip the assertion.
|
|
||||||
func TestNoGoroutineLeaks(t *testing.T) {
|
|
||||||
defer goleak.VerifyNone(t)
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
r := router.NewR(t, myControl, theirControl)
|
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
r.RenderFlow()
|
|
||||||
|
|
||||||
// Settle period: Stop() is non-blocking; the wg-driven goroutines need
|
|
||||||
// a moment to drain. goleak retries internally too, but a short explicit
|
|
||||||
// settle reduces flakes when the suite is busy.
|
|
||||||
time.Sleep(50 * time.Millisecond)
|
|
||||||
}
|
|
||||||
@@ -1,225 +0,0 @@
|
|||||||
//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()
|
|
||||||
}
|
|
||||||
+82
-318
@@ -13,7 +13,6 @@ import (
|
|||||||
"regexp"
|
"regexp"
|
||||||
"sort"
|
"sort"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -25,19 +24,6 @@ import (
|
|||||||
"golang.org/x/exp/maps"
|
"golang.org/x/exp/maps"
|
||||||
)
|
)
|
||||||
|
|
||||||
// outNatKey is the (from, to) pair used by outNat. Comparable struct, so it works as a map key without the
|
|
||||||
// allocation cost of a string-concat key.
|
|
||||||
type outNatKey struct {
|
|
||||||
from, to netip.AddrPort
|
|
||||||
}
|
|
||||||
|
|
||||||
// fannedPacket pairs a UDP TX packet with its source control so the router can route it after popping from
|
|
||||||
// the fan-in channel.
|
|
||||||
type fannedPacket struct {
|
|
||||||
from *nebula.Control
|
|
||||||
pkt *udp.Packet
|
|
||||||
}
|
|
||||||
|
|
||||||
type R struct {
|
type R struct {
|
||||||
// Simple map of the ip:port registered on a control to the control
|
// Simple map of the ip:port registered on a control to the control
|
||||||
// Basically a router, right?
|
// Basically a router, right?
|
||||||
@@ -48,28 +34,12 @@ type R struct {
|
|||||||
|
|
||||||
// A last used map, if an inbound packet hit the inNat map then
|
// A last used map, if an inbound packet hit the inNat map then
|
||||||
// all return packets should use the same last used inbound address for the outbound sender
|
// all return packets should use the same last used inbound address for the outbound sender
|
||||||
outNat map[outNatKey]netip.AddrPort
|
// map[from address + ":" + to address] => ip:port to rewrite in the udp packet to receiver
|
||||||
|
outNat map[string]netip.AddrPort
|
||||||
|
|
||||||
// A map of vpn ip to the nebula control it belongs to
|
// A map of vpn ip to the nebula control it belongs to
|
||||||
vpnControls map[netip.Addr]*nebula.Control
|
vpnControls map[netip.Addr]*nebula.Control
|
||||||
|
|
||||||
// Cached select infrastructure for RouteForAllUntilTxTun.
|
|
||||||
// The controls map is immutable after NewR so the cases are good for the test lifetime.
|
|
||||||
// We only rebuild if a different receiver is asked.
|
|
||||||
selRecvCtl *nebula.Control
|
|
||||||
selCases []reflect.SelectCase
|
|
||||||
selCtls []*nebula.Control
|
|
||||||
|
|
||||||
// Optional fan-in mode for hot-path benchmarks: one forwarder goroutine per control drains UDP TX into udpFanIn,
|
|
||||||
// so RouteForAllUntilTxTun can do a fixed 2-way native select instead of paying reflect.Select per call.
|
|
||||||
// Off by default (would otherwise interleave with tests that use GetFromUDP directly on the same control).
|
|
||||||
// Enabled by EnableFanIn.
|
|
||||||
udpFanIn chan fannedPacket
|
|
||||||
stopFanIn chan struct{}
|
|
||||||
fanInWG sync.WaitGroup
|
|
||||||
fanInMu sync.Mutex
|
|
||||||
fanInOn atomic.Bool
|
|
||||||
|
|
||||||
ignoreFlows []ignoreFlow
|
ignoreFlows []ignoreFlow
|
||||||
flow []flowEntry
|
flow []flowEntry
|
||||||
|
|
||||||
@@ -114,28 +84,6 @@ 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() {
|
||||||
@@ -171,7 +119,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
|||||||
controls: make(map[netip.AddrPort]*nebula.Control),
|
controls: make(map[netip.AddrPort]*nebula.Control),
|
||||||
vpnControls: make(map[netip.Addr]*nebula.Control),
|
vpnControls: make(map[netip.Addr]*nebula.Control),
|
||||||
inNat: make(map[netip.AddrPort]*nebula.Control),
|
inNat: make(map[netip.AddrPort]*nebula.Control),
|
||||||
outNat: make(map[outNatKey]netip.AddrPort),
|
outNat: make(map[string]netip.AddrPort),
|
||||||
flow: []flowEntry{},
|
flow: []flowEntry{},
|
||||||
ignoreFlows: []ignoreFlow{},
|
ignoreFlows: []ignoreFlow{},
|
||||||
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
||||||
@@ -205,10 +153,8 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
|||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return
|
||||||
case <-clockSource.C:
|
case <-clockSource.C:
|
||||||
r.Lock()
|
|
||||||
r.renderHostmaps("clock tick")
|
r.renderHostmaps("clock tick")
|
||||||
r.renderFlow()
|
r.renderFlow()
|
||||||
r.Unlock()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
@@ -234,21 +180,15 @@ func (r *R) AddRoute(ip netip.Addr, port uint16, c *nebula.Control) {
|
|||||||
// RenderFlow renders the packet flow seen up until now and stops further automatic renders from happening.
|
// RenderFlow renders the packet flow seen up until now and stops further automatic renders from happening.
|
||||||
func (r *R) RenderFlow() {
|
func (r *R) RenderFlow() {
|
||||||
r.cancelRender()
|
r.cancelRender()
|
||||||
r.Lock()
|
|
||||||
defer r.Unlock()
|
|
||||||
r.renderFlow()
|
r.renderFlow()
|
||||||
}
|
}
|
||||||
|
|
||||||
// CancelFlowLogs stops flow logs from being tracked and destroys any logs already collected
|
// CancelFlowLogs stops flow logs from being tracked and destroys any logs already collected
|
||||||
func (r *R) CancelFlowLogs() {
|
func (r *R) CancelFlowLogs() {
|
||||||
r.cancelRender()
|
r.cancelRender()
|
||||||
r.Lock()
|
|
||||||
r.flow = nil
|
r.flow = nil
|
||||||
r.Unlock()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// renderFlow writes the flow log to disk. Caller must hold r.Lock. renderFlow reads r.flow / r.additionalGraphs and
|
|
||||||
// the *packet pointers stashed inside, all of which are mutated under the same lock by routing paths.
|
|
||||||
func (r *R) renderFlow() {
|
func (r *R) renderFlow() {
|
||||||
if r.flow == nil {
|
if r.flow == nil {
|
||||||
return
|
return
|
||||||
@@ -271,7 +211,7 @@ func (r *R) renderFlow() {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
addr := e.packet.fromAddr()
|
addr := e.packet.from.GetUDPAddr()
|
||||||
if _, ok := participants[addr]; ok {
|
if _, ok := participants[addr]; ok {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -290,6 +230,7 @@ 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)
|
||||||
@@ -301,22 +242,21 @@ 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 = "->>"
|
||||||
}
|
}
|
||||||
|
|
||||||
detail := fmt.Sprintf("%s(%s), index %v, counter: %v",
|
fmt.Fprintf(f,
|
||||||
p.h.TypeName(), p.h.SubTypeName(), p.h.RemoteIndex, p.h.MessageCounter)
|
" %s%s%s: %s(%s), index %v, counter: %v\n",
|
||||||
if p.parseErr != nil {
|
normalizeName(p.from.GetUDPAddr().String()),
|
||||||
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.toAddr().String()),
|
normalizeName(p.to.GetUDPAddr().String()),
|
||||||
detail,
|
h.TypeName(), h.SubTypeName(), h.RemoteIndex, h.MessageCounter,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -430,34 +370,29 @@ 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)))
|
||||||
|
|
||||||
var h header.H
|
if len(r.ignoreFlows) > 0 {
|
||||||
var parseErr error
|
var h header.H
|
||||||
if !tun {
|
err := h.Parse(p.Data)
|
||||||
parseErr = h.Parse(p.Data)
|
if err != nil {
|
||||||
}
|
panic(err)
|
||||||
|
|
||||||
// Decide before copying, the copy comes from a freelist and an ignored packet would never be released
|
|
||||||
for _, i := range r.ignoreFlows {
|
|
||||||
if tun {
|
|
||||||
if i.tun.HasValue && i.tun.IsTrue {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// A packet we could not parse has no type to match against, so no rule can ignore it
|
for _, i := range r.ignoreFlows {
|
||||||
if parseErr == nil && i.messageType == h.Type && i.subType == h.Subtype {
|
if !tun {
|
||||||
return nil
|
if i.messageType == h.Type && i.subType == h.Subtype {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
} else if i.tun.HasValue && i.tun.IsTrue {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fp := &packet{
|
fp := &packet{
|
||||||
from: from,
|
from: from,
|
||||||
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})
|
||||||
@@ -499,157 +434,68 @@ func (r *R) RouteUntilTxTun(sender *nebula.Control, receiver *nebula.Control) []
|
|||||||
panic("No control for udp tx " + a.String())
|
panic("No control for udp tx " + a.String())
|
||||||
}
|
}
|
||||||
fp := r.unlockedInjectFlow(sender, c, p, false)
|
fp := r.unlockedInjectFlow(sender, c, p, false)
|
||||||
c.InjectUDPPacket(p) // copies internally; original is ours to release
|
c.InjectUDPPacket(p)
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// RouteForAllUntilTxTun will route for everyone and return when a packet is seen on the receiver's tun.
|
// RouteForAllUntilTxTun will route for everyone and return when a packet is seen on receivers tun
|
||||||
// If a control's UDP TX address can't be matched to a registered control, we panic.
|
// If the router doesn't have the nebula controller for that address, we panic
|
||||||
//
|
|
||||||
// For allocation-sensitive callers (hot-path benchmarks, in particular relay
|
|
||||||
// benches with 3+ controls), call EnableFanIn() first.
|
|
||||||
func (r *R) RouteForAllUntilTxTun(receiver *nebula.Control) []byte {
|
func (r *R) RouteForAllUntilTxTun(receiver *nebula.Control) []byte {
|
||||||
if r.fanInOn.Load() {
|
|
||||||
return r.routeFanIn(receiver)
|
|
||||||
}
|
|
||||||
return r.routeReflect(receiver)
|
|
||||||
}
|
|
||||||
|
|
||||||
// routeFanIn is the alloc-free path used when EnableFanIn is in effect.
|
|
||||||
func (r *R) routeFanIn(receiver *nebula.Control) []byte {
|
|
||||||
tunTx := receiver.GetTunTxChan()
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case p := <-tunTx:
|
|
||||||
r.Lock()
|
|
||||||
if r.flow != nil {
|
|
||||||
np := udp.Packet{Data: make([]byte, len(p))}
|
|
||||||
copy(np.Data, p)
|
|
||||||
r.unlockedInjectFlow(receiver, receiver, &np, true)
|
|
||||||
}
|
|
||||||
r.Unlock()
|
|
||||||
return p
|
|
||||||
case fp := <-r.udpFanIn:
|
|
||||||
r.routeUDP(fp.from, fp.pkt)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// routeReflect is the default reflect.Select-based path. Pays the boxing allocation per call but doesn't interfere
|
|
||||||
// with tests that pull packets directly from controls' UDP TX channels via GetFromUDP.
|
|
||||||
func (r *R) routeReflect(receiver *nebula.Control) []byte {
|
|
||||||
sc, cm := r.selectCasesFor(receiver)
|
|
||||||
for {
|
|
||||||
x, rx, _ := reflect.Select(sc)
|
|
||||||
if x == 0 {
|
|
||||||
p := rx.Interface().([]byte)
|
|
||||||
r.Lock()
|
|
||||||
if r.flow != nil {
|
|
||||||
np := udp.Packet{Data: make([]byte, len(p))}
|
|
||||||
copy(np.Data, p)
|
|
||||||
r.unlockedInjectFlow(cm[x], cm[x], &np, true)
|
|
||||||
}
|
|
||||||
r.Unlock()
|
|
||||||
return p
|
|
||||||
}
|
|
||||||
r.routeUDP(cm[x], rx.Interface().(*udp.Packet))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// EnableFanIn switches RouteForAllUntilTxTun to the alloc-free fan-in path.
|
|
||||||
// One forwarder goroutine per registered control drains UDP TX into a shared channel that RouteForAllUntilTxTun selects
|
|
||||||
// on alongside the receiver's TUN TX channel.
|
|
||||||
func (r *R) EnableFanIn() {
|
|
||||||
r.fanInMu.Lock()
|
|
||||||
defer r.fanInMu.Unlock()
|
|
||||||
if r.fanInOn.Load() {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
r.udpFanIn = make(chan fannedPacket, 32)
|
|
||||||
r.stopFanIn = make(chan struct{})
|
|
||||||
for _, c := range r.controls {
|
|
||||||
r.startFanInWorker(c)
|
|
||||||
}
|
|
||||||
r.fanInOn.Store(true)
|
|
||||||
r.t.Cleanup(r.stopFanInWorkers)
|
|
||||||
}
|
|
||||||
|
|
||||||
// startFanInWorker spawns a goroutine that drains c's UDP TX into r.udpFanIn.
|
|
||||||
func (r *R) startFanInWorker(c *nebula.Control) {
|
|
||||||
r.fanInWG.Add(1)
|
|
||||||
udpTx := c.GetUDPTxChan()
|
|
||||||
go func() {
|
|
||||||
defer r.fanInWG.Done()
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-r.stopFanIn:
|
|
||||||
return
|
|
||||||
case p := <-udpTx:
|
|
||||||
select {
|
|
||||||
case <-r.stopFanIn:
|
|
||||||
p.Release()
|
|
||||||
return
|
|
||||||
case r.udpFanIn <- fannedPacket{from: c, pkt: p}:
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
}
|
|
||||||
|
|
||||||
// stopFanInWorkers signals the fan-in goroutines to exit and waits for them.
|
|
||||||
func (r *R) stopFanInWorkers() {
|
|
||||||
r.fanInMu.Lock()
|
|
||||||
wasOn := r.fanInOn.Swap(false)
|
|
||||||
r.fanInMu.Unlock()
|
|
||||||
if !wasOn {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
close(r.stopFanIn)
|
|
||||||
r.fanInWG.Wait()
|
|
||||||
}
|
|
||||||
|
|
||||||
// routeUDP forwards a UDP TX packet from the named source control to the destination control derived from p.To,
|
|
||||||
// releasing the source packet after InjectUDPPacket has copied its bytes into a fresh pool slot.
|
|
||||||
func (r *R) routeUDP(from *nebula.Control, p *udp.Packet) {
|
|
||||||
r.Lock()
|
|
||||||
defer r.Unlock()
|
|
||||||
a := from.GetUDPAddr()
|
|
||||||
c := r.getControl(a, p.To, p)
|
|
||||||
if c == nil {
|
|
||||||
panic(fmt.Sprintf("No control for udp tx %s", p.To))
|
|
||||||
}
|
|
||||||
fp := r.unlockedInjectFlow(from, c, p, false)
|
|
||||||
c.InjectUDPPacket(p) // copies internally; original is ours to release
|
|
||||||
fp.WasReceived()
|
|
||||||
p.Release()
|
|
||||||
}
|
|
||||||
|
|
||||||
// selectCasesFor returns the SelectCase array used by routeReflect: one slot for the receiver's TUN TX channel followed
|
|
||||||
// by one per control's UDP TX channel. Cached for the test lifetime, only rebuilt if the receiver changes.
|
|
||||||
func (r *R) selectCasesFor(receiver *nebula.Control) ([]reflect.SelectCase, []*nebula.Control) {
|
|
||||||
r.Lock()
|
|
||||||
defer r.Unlock()
|
|
||||||
if r.selRecvCtl == receiver && r.selCases != nil {
|
|
||||||
return r.selCases, r.selCtls
|
|
||||||
}
|
|
||||||
sc := make([]reflect.SelectCase, len(r.controls)+1)
|
sc := make([]reflect.SelectCase, len(r.controls)+1)
|
||||||
cm := make([]*nebula.Control, len(r.controls)+1)
|
cm := make([]*nebula.Control, len(r.controls)+1)
|
||||||
sc[0] = reflect.SelectCase{Dir: reflect.SelectRecv, Chan: reflect.ValueOf(receiver.GetTunTxChan())}
|
|
||||||
cm[0] = receiver
|
i := 0
|
||||||
i := 1
|
sc[i] = reflect.SelectCase{
|
||||||
|
Dir: reflect.SelectRecv,
|
||||||
|
Chan: reflect.ValueOf(receiver.GetTunTxChan()),
|
||||||
|
Send: reflect.Value{},
|
||||||
|
}
|
||||||
|
cm[i] = receiver
|
||||||
|
|
||||||
|
i++
|
||||||
for _, c := range r.controls {
|
for _, c := range r.controls {
|
||||||
sc[i] = reflect.SelectCase{Dir: reflect.SelectRecv, Chan: reflect.ValueOf(c.GetUDPTxChan())}
|
sc[i] = reflect.SelectCase{
|
||||||
|
Dir: reflect.SelectRecv,
|
||||||
|
Chan: reflect.ValueOf(c.GetUDPTxChan()),
|
||||||
|
Send: reflect.Value{},
|
||||||
|
}
|
||||||
|
|
||||||
cm[i] = c
|
cm[i] = c
|
||||||
i++
|
i++
|
||||||
}
|
}
|
||||||
r.selRecvCtl = receiver
|
|
||||||
r.selCases = sc
|
for {
|
||||||
r.selCtls = cm
|
x, rx, _ := reflect.Select(sc)
|
||||||
return sc, cm
|
r.Lock()
|
||||||
|
|
||||||
|
if x == 0 {
|
||||||
|
// we are the tun tx, we can exit
|
||||||
|
p := rx.Interface().([]byte)
|
||||||
|
np := udp.Packet{Data: make([]byte, len(p))}
|
||||||
|
copy(np.Data, p)
|
||||||
|
|
||||||
|
r.unlockedInjectFlow(cm[x], cm[x], &np, true)
|
||||||
|
r.Unlock()
|
||||||
|
return p
|
||||||
|
|
||||||
|
} else {
|
||||||
|
// we are a udp tx, route and continue
|
||||||
|
p := rx.Interface().(*udp.Packet)
|
||||||
|
a := cm[x].GetUDPAddr()
|
||||||
|
c := r.getControl(a, p.To, p)
|
||||||
|
if c == nil {
|
||||||
|
r.Unlock()
|
||||||
|
panic(fmt.Sprintf("No control for udp tx %s", p.To))
|
||||||
|
}
|
||||||
|
fp := r.unlockedInjectFlow(cm[x], c, p, false)
|
||||||
|
c.InjectUDPPacket(p)
|
||||||
|
fp.WasReceived()
|
||||||
|
}
|
||||||
|
r.Unlock()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// RouteExitFunc will call the whatDo func with each udp packet from sender.
|
// RouteExitFunc will call the whatDo func with each udp packet from sender.
|
||||||
@@ -676,7 +522,6 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
switch e {
|
switch e {
|
||||||
case ExitNow:
|
case ExitNow:
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
return
|
return
|
||||||
|
|
||||||
case RouteAndExit:
|
case RouteAndExit:
|
||||||
@@ -684,7 +529,6 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
return
|
return
|
||||||
|
|
||||||
case KeepRouting:
|
case KeepRouting:
|
||||||
@@ -697,7 +541,6 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -717,81 +560,6 @@ 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 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 {
|
||||||
@@ -873,7 +641,6 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
switch e {
|
switch e {
|
||||||
case ExitNow:
|
case ExitNow:
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
return
|
return
|
||||||
|
|
||||||
case RouteAndExit:
|
case RouteAndExit:
|
||||||
@@ -881,7 +648,6 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
return
|
return
|
||||||
|
|
||||||
case KeepRouting:
|
case KeepRouting:
|
||||||
@@ -893,7 +659,6 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
||||||
}
|
}
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -937,20 +702,19 @@ func (r *R) FlushAll() {
|
|||||||
}
|
}
|
||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// getControl performs or seeds NAT translation and returns the control for toAddr, p from fields may change
|
// getControl performs or seeds NAT translation and returns the control for toAddr, p from fields may change
|
||||||
// This is an internal router function, the caller must hold the lock
|
// This is an internal router function, the caller must hold the lock
|
||||||
func (r *R) getControl(fromAddr, toAddr netip.AddrPort, p *udp.Packet) *nebula.Control {
|
func (r *R) getControl(fromAddr, toAddr netip.AddrPort, p *udp.Packet) *nebula.Control {
|
||||||
if newAddr, ok := r.outNat[outNatKey{from: fromAddr, to: toAddr}]; ok {
|
if newAddr, ok := r.outNat[fromAddr.String()+":"+toAddr.String()]; ok {
|
||||||
p.From = newAddr
|
p.From = newAddr
|
||||||
}
|
}
|
||||||
|
|
||||||
c, ok := r.inNat[toAddr]
|
c, ok := r.inNat[toAddr]
|
||||||
if ok {
|
if ok {
|
||||||
r.outNat[outNatKey{from: c.GetUDPAddr(), to: fromAddr}] = toAddr
|
r.outNat[c.GetUDPAddr().String()+":"+fromAddr.String()] = toAddr
|
||||||
return c
|
return c
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,125 +0,0 @@
|
|||||||
//go:build e2e_testing
|
|
||||||
// +build e2e_testing
|
|
||||||
|
|
||||||
package e2e
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/ed25519"
|
|
||||||
"crypto/rand"
|
|
||||||
"encoding/pem"
|
|
||||||
"net"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/cert_test"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"golang.org/x/crypto/ssh"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestSSHDLifecycle(t *testing.T) {
|
|
||||||
// TestSSHDLifecycle exercises the in-process sshd through several config reloads and a Control.Stop.
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(
|
|
||||||
cert.Version1, cert.Curve_CURVE25519,
|
|
||||||
time.Now(), time.Now().Add(10*time.Minute),
|
|
||||||
nil, nil, []string{},
|
|
||||||
)
|
|
||||||
|
|
||||||
hostKeyPEM := generateSSHHostKey(t)
|
|
||||||
clientSigner, clientAuthKey := generateSSHClientKey(t)
|
|
||||||
sshdAddr := allocLoopbackPort(t)
|
|
||||||
|
|
||||||
overrides := m{
|
|
||||||
"sshd": m{
|
|
||||||
"enabled": true,
|
|
||||||
"listen": sshdAddr,
|
|
||||||
"host_key": hostKeyPEM,
|
|
||||||
"authorized_users": []m{{
|
|
||||||
"user": "tester",
|
|
||||||
"keys": []string{clientAuthKey},
|
|
||||||
}},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
control, _, _, _ := newSimpleServer(cert.Version1, ca, caKey, "sshd-test", "10.222.0.1/24", overrides)
|
|
||||||
control.Start()
|
|
||||||
t.Cleanup(func() { control.Stop() })
|
|
||||||
|
|
||||||
// sshd binds in a goroutine after Start returns; wait for it.
|
|
||||||
require.Eventually(t, func() bool { return canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
|
||||||
"sshd never started listening")
|
|
||||||
|
|
||||||
for i := 1; i <= 3; i++ {
|
|
||||||
out := sshExecReload(t, sshdAddr, clientSigner)
|
|
||||||
assert.Contains(t, out, "Reloading config", "reload cycle %d", i)
|
|
||||||
require.Eventually(t, func() bool { return canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
|
||||||
"sshd not listening after reload cycle %d", i)
|
|
||||||
}
|
|
||||||
|
|
||||||
control.Stop()
|
|
||||||
require.Eventually(t, func() bool { return !canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
|
||||||
"sshd still listening after Control.Stop")
|
|
||||||
}
|
|
||||||
|
|
||||||
func canDial(addr string) bool {
|
|
||||||
c, err := net.DialTimeout("tcp", addr, 100*time.Millisecond)
|
|
||||||
if err != nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
_ = c.Close()
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// allocLoopbackPort grabs an unused TCP port on 127.0.0.1, closes it, and returns the address. There
|
|
||||||
// is a small race between releasing the port and the sshd reclaiming it; in practice the OS keeps the
|
|
||||||
// port available long enough for the test to bind it.
|
|
||||||
func allocLoopbackPort(t *testing.T) string {
|
|
||||||
t.Helper()
|
|
||||||
l, err := net.Listen("tcp", "127.0.0.1:0")
|
|
||||||
require.NoError(t, err)
|
|
||||||
addr := l.Addr().String()
|
|
||||||
require.NoError(t, l.Close())
|
|
||||||
return addr
|
|
||||||
}
|
|
||||||
|
|
||||||
func generateSSHHostKey(t *testing.T) string {
|
|
||||||
t.Helper()
|
|
||||||
_, priv, err := ed25519.GenerateKey(rand.Reader)
|
|
||||||
require.NoError(t, err)
|
|
||||||
block, err := ssh.MarshalPrivateKey(priv, "nebula-e2e-host")
|
|
||||||
require.NoError(t, err)
|
|
||||||
return string(pem.EncodeToMemory(block))
|
|
||||||
}
|
|
||||||
|
|
||||||
func generateSSHClientKey(t *testing.T) (ssh.Signer, string) {
|
|
||||||
t.Helper()
|
|
||||||
_, priv, err := ed25519.GenerateKey(rand.Reader)
|
|
||||||
require.NoError(t, err)
|
|
||||||
signer, err := ssh.NewSignerFromKey(priv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
auth := strings.TrimSpace(string(ssh.MarshalAuthorizedKey(signer.PublicKey())))
|
|
||||||
return signer, auth
|
|
||||||
}
|
|
||||||
|
|
||||||
func sshExecReload(t *testing.T, addr string, signer ssh.Signer) string {
|
|
||||||
t.Helper()
|
|
||||||
cfg := &ssh.ClientConfig{
|
|
||||||
User: "tester",
|
|
||||||
Auth: []ssh.AuthMethod{ssh.PublicKeys(signer)},
|
|
||||||
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
|
|
||||||
Timeout: 2 * time.Second,
|
|
||||||
}
|
|
||||||
client, err := ssh.Dial("tcp", addr, cfg)
|
|
||||||
require.NoError(t, err)
|
|
||||||
defer client.Close()
|
|
||||||
|
|
||||||
sess, err := client.NewSession()
|
|
||||||
require.NoError(t, err)
|
|
||||||
defer sess.Close()
|
|
||||||
|
|
||||||
// reload tears the channel down before sending exit-status, so Output returns an error on the
|
|
||||||
// channel close. The output buffer still has whatever the reload callback wrote before that.
|
|
||||||
out, _ := sess.Output("reload")
|
|
||||||
return string(out)
|
|
||||||
}
|
|
||||||
+8
-109
@@ -15,12 +15,10 @@ 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"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestDropInactiveTunnels(t *testing.T) {
|
func TestDropInactiveTunnels(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -43,8 +41,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 := myControl.GetHostmapIndexCount()
|
myIndexes := len(myControl.GetHostmap().Indexes)
|
||||||
theirIndexes := theirControl.GetHostmapIndexCount()
|
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
||||||
if myIndexes == 0 && theirIndexes == 0 {
|
if myIndexes == 0 && theirIndexes == 0 {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
@@ -65,7 +63,6 @@ func TestDropInactiveTunnels(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCertUpgrade(t *testing.T) {
|
func TestCertUpgrade(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -160,7 +157,6 @@ func TestCertUpgrade(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCertDowngrade(t *testing.T) {
|
func TestCertDowngrade(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -259,7 +255,6 @@ func TestCertDowngrade(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCertMismatchCorrection(t *testing.T) {
|
func TestCertMismatchCorrection(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -327,7 +322,6 @@ func TestCertMismatchCorrection(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCrossStackRelaysWork(t *testing.T) {
|
func TestCrossStackRelaysWork(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24,fc00::1/64", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24,fc00::1/64", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "10.128.0.128/24,fc00::128/64", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "10.128.0.128/24,fc00::128/64", m{"relay": m{"am_relay": true}})
|
||||||
@@ -356,14 +350,14 @@ func TestCrossStackRelaysWork(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80)
|
||||||
|
|
||||||
t.Log("reply?")
|
t.Log("reply?")
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them"))
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnV6.Addr(), myVpnV6.Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnV6.Addr(), myVpnV6.Addr(), 80, 80)
|
||||||
|
|
||||||
@@ -374,102 +368,7 @@ 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()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", m{"tunnels": m{"drop_inactive": true, "inactivity_timeout": "5s"}})
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", m{"tunnels": m{"drop_inactive": true, "inactivity_timeout": "5s"}})
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", m{"tunnels": m{"drop_inactive": true, "inactivity_timeout": "10m"}})
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", m{"tunnels": m{"drop_inactive": true, "inactivity_timeout": "10m"}})
|
||||||
@@ -493,8 +392,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) {
|
|||||||
|
|
||||||
waitStart := time.Now()
|
waitStart := time.Now()
|
||||||
for {
|
for {
|
||||||
myIndexes := myControl.GetHostmapIndexCount()
|
myIndexes := len(myControl.GetHostmap().Indexes)
|
||||||
theirIndexes := theirControl.GetHostmapIndexCount()
|
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
||||||
if myIndexes == 0 && theirIndexes == 0 {
|
if myIndexes == 0 && theirIndexes == 0 {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
@@ -548,8 +447,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 := myControl.GetHostmapIndexCount()
|
myIndexes := len(myControl.GetHostmap().Indexes)
|
||||||
theirIndexes := theirControl.GetHostmapIndexCount()
|
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
||||||
if myIndexes == 0 {
|
if myIndexes == 0 {
|
||||||
t.Fatal("myIndexes should not be 0")
|
t.Fatal("myIndexes should not be 0")
|
||||||
}
|
}
|
||||||
|
|||||||
+15
-52
@@ -138,22 +138,6 @@ listen:
|
|||||||
# max, net.core.rmem_max and net.core.wmem_max
|
# max, net.core.rmem_max and net.core.wmem_max
|
||||||
#read_buffer: 10485760
|
#read_buffer: 10485760
|
||||||
#write_buffer: 10485760
|
#write_buffer: 10485760
|
||||||
|
|
||||||
# On Windows only
|
|
||||||
# When true, Nebula installs a WFP (Windows Filtering Platform) PERMIT filter scoped to UDP at the listener port.
|
|
||||||
# WFP sits below Windows Defender Firewall, so this lets peer handshakes reach Nebula's outside socket regardless
|
|
||||||
# of WDF's inbound rules.
|
|
||||||
# Default true; set to false to leave WDF in charge of inbound decisions on the listener port. Not reloadable.
|
|
||||||
#windows_bypass_wdf: true
|
|
||||||
|
|
||||||
# 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.
|
||||||
@@ -179,21 +163,17 @@ listen:
|
|||||||
|
|
||||||
punchy:
|
punchy:
|
||||||
# Continues to punch inbound/outbound at a regular interval to avoid expiration of firewall nat mappings
|
# Continues to punch inbound/outbound at a regular interval to avoid expiration of firewall nat mappings
|
||||||
# This setting is reloadable.
|
|
||||||
punch: true
|
punch: true
|
||||||
|
|
||||||
# respond means that a node you are trying to reach will connect back out to you if your hole punching fails
|
# respond means that a node you are trying to reach will connect back out to you if your hole punching fails
|
||||||
# this is extremely useful if one node is behind a difficult nat, such as a symmetric NAT
|
# this is extremely useful if one node is behind a difficult nat, such as a symmetric NAT
|
||||||
# Default is false
|
# Default is false
|
||||||
# This setting is reloadable.
|
|
||||||
#respond: true
|
#respond: true
|
||||||
|
|
||||||
# delays a punch response for misbehaving NATs, default is 1 second.
|
# delays a punch response for misbehaving NATs, default is 1 second.
|
||||||
# This setting is reloadable.
|
|
||||||
#delay: 1s
|
#delay: 1s
|
||||||
|
|
||||||
# set the delay before attempting punchy.respond. Default is 5 seconds. respond must be true to take effect.
|
# set the delay before attempting punchy.respond. Default is 5 seconds. respond must be true to take effect.
|
||||||
# This setting is reloadable.
|
|
||||||
#respond_delay: 5s
|
#respond_delay: 5s
|
||||||
|
|
||||||
# Cipher allows you to choose between the available ciphers for your network. Options are chachapoly or aes
|
# Cipher allows you to choose between the available ciphers for your network. Options are chachapoly or aes
|
||||||
@@ -302,24 +282,6 @@ tun:
|
|||||||
# metric: 100
|
# metric: 100
|
||||||
# install: true
|
# install: true
|
||||||
|
|
||||||
# On Windows only, sets the network category of the nebula interface. Without this, Windows often
|
|
||||||
# leaves the network as "Unidentified" and treats it as Public, which makes the host firewall more
|
|
||||||
# restrictive than you usually want for an overlay between trusted peers. Valid values:
|
|
||||||
# private - treat the nebula network as a private/trusted network (default)
|
|
||||||
# public - treat it as a public/untrusted network
|
|
||||||
# domain - treat it as a domain-authenticated network
|
|
||||||
# unset - leave whatever Windows decided alone
|
|
||||||
# Not reloadable.
|
|
||||||
#network_category: private
|
|
||||||
|
|
||||||
# On Windows only
|
|
||||||
# When true, Nebula installs a WFP (Windows Filtering Platform) PERMIT filter scoped to the nebula adapter LUID.
|
|
||||||
# WFP sits below Windows Defender Firewall, so this lets inbound traffic through regardless of WDF rules.
|
|
||||||
# Filters are auto-removed when the adapter goes away.
|
|
||||||
# See listen.windows_bypass_wdf for the matching control over inbound to nebula's outside UDP listener.
|
|
||||||
# Default true; set to false to leave WDF in charge of inbound decisions on the nebula interface. Not reloadable.
|
|
||||||
#windows_bypass_wdf: true
|
|
||||||
|
|
||||||
# On linux only, set to true to manage unsafe routes directly on the system route table with gateway routes instead of
|
# On linux only, set to true to manage unsafe routes directly on the system route table with gateway routes instead of
|
||||||
# in nebula configuration files. Default false, not reloadable.
|
# in nebula configuration files. Default false, not reloadable.
|
||||||
#use_system_route_table: false
|
#use_system_route_table: false
|
||||||
@@ -330,21 +292,24 @@ tun:
|
|||||||
|
|
||||||
# Configure logging level
|
# Configure logging level
|
||||||
logging:
|
logging:
|
||||||
# trace, debug, info, warn, or error. Default is info and is reloadable.
|
# panic, fatal, error, warning, info, or debug. Default is info and is reloadable.
|
||||||
# fatal and panic are accepted for backwards compatibility and map to error.
|
#NOTE: Debug mode can log remotely controlled/untrusted data which can quickly fill a disk in some
|
||||||
#NOTE: Debug and trace modes can log remotely controlled/untrusted data which can quickly fill a disk in some
|
# scenarios. Debug logging is also CPU intensive and will decrease performance overall.
|
||||||
# scenarios. Debug and trace logging are also CPU intensive and will decrease performance overall.
|
# Only enable debug logging while actively investigating an issue.
|
||||||
# Only enable debug or trace logging while actively investigating an issue.
|
|
||||||
level: info
|
level: info
|
||||||
# json or text formats currently available. Default is text.
|
# json or text formats currently available. Default is text
|
||||||
format: text
|
format: text
|
||||||
# Disable timestamp logging. Useful when output is redirected to a logging system that already adds timestamps. Default is false.
|
# Disable timestamp logging. useful when output is redirected to logging system that already adds timestamps. Default is false
|
||||||
#disable_timestamp: true
|
#disable_timestamp: true
|
||||||
# Timestamps use RFC3339Nano ("2006-01-02T15:04:05.999999999Z07:00") and are not configurable.
|
# timestamp format is specified in Go time format, see:
|
||||||
|
# https://golang.org/pkg/time/#pkg-constants
|
||||||
|
# default when `format: json`: "2006-01-02T15:04:05Z07:00" (RFC3339)
|
||||||
|
# default when `format: text`:
|
||||||
|
# when TTY attached: seconds since beginning of execution
|
||||||
|
# otherwise: "2006-01-02T15:04:05Z07:00" (RFC3339)
|
||||||
|
# As an example, to log as RFC3339 with millisecond precision, set to:
|
||||||
|
#timestamp_format: "2006-01-02T15:04:05.000Z07:00"
|
||||||
|
|
||||||
# The stats section is reloadable. A HUP may change the backend, toggle stats
|
|
||||||
# on or off, switch the listen/host address, or pick up new DNS for the
|
|
||||||
# configured graphite host.
|
|
||||||
#stats:
|
#stats:
|
||||||
#type: graphite
|
#type: graphite
|
||||||
#prefix: nebula
|
#prefix: nebula
|
||||||
@@ -362,12 +327,10 @@ logging:
|
|||||||
# enables counter metrics for meta packets
|
# enables counter metrics for meta packets
|
||||||
# e.g.: `messages.tx.handshake`
|
# e.g.: `messages.tx.handshake`
|
||||||
# NOTE: `message.{tx,rx}.recv_error` is always emitted
|
# NOTE: `message.{tx,rx}.recv_error` is always emitted
|
||||||
# Not reloadable.
|
|
||||||
#message_metrics: false
|
#message_metrics: false
|
||||||
|
|
||||||
# enables detailed counter metrics for lighthouse packets
|
# enables detailed counter metrics for lighthouse packets
|
||||||
# e.g.: `lighthouse.rx.HostQuery`
|
# e.g.: `lighthouse.rx.HostQuery`
|
||||||
# Not reloadable.
|
|
||||||
#lighthouse_metrics: false
|
#lighthouse_metrics: false
|
||||||
|
|
||||||
# Handshake Manager Settings
|
# Handshake Manager Settings
|
||||||
@@ -405,7 +368,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 "Destination unreachable: Communication administratively prohibited" packet.
|
# - For other protocols, this will be an ICMP port unreachable packet.
|
||||||
outbound_action: drop
|
outbound_action: drop
|
||||||
inbound_action: drop
|
inbound_action: drop
|
||||||
|
|
||||||
|
|||||||
@@ -7,9 +7,9 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"os"
|
"os"
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
"github.com/slackhq/nebula/service"
|
"github.com/slackhq/nebula/service"
|
||||||
)
|
)
|
||||||
@@ -64,7 +64,8 @@ pki:
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
logger := logging.NewLogger(os.Stdout)
|
logger := logrus.New()
|
||||||
|
logger.Out = os.Stdout
|
||||||
|
|
||||||
ctrl, err := nebula.Main(&cfg, false, "custom-app", logger, overlay.NewUserDeviceFromConfig)
|
ctrl, err := nebula.Main(&cfg, false, "custom-app", logger, overlay.NewUserDeviceFromConfig)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -8,15 +8,6 @@ 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
|
||||||
|
|||||||
+155
-98
@@ -1,13 +1,11 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"hash/fnv"
|
"hash/fnv"
|
||||||
"log/slog"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"reflect"
|
"reflect"
|
||||||
"slices"
|
"slices"
|
||||||
@@ -18,9 +16,11 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/cert"
|
"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/firewall/events"
|
||||||
)
|
)
|
||||||
|
|
||||||
type FirewallInterface interface {
|
type FirewallInterface interface {
|
||||||
@@ -44,8 +44,8 @@ type Firewall struct {
|
|||||||
InRules *FirewallTable
|
InRules *FirewallTable
|
||||||
OutRules *FirewallTable
|
OutRules *FirewallTable
|
||||||
|
|
||||||
InboundSendReject bool
|
InSendReject bool
|
||||||
OutboundSendReject bool
|
OutSendReject bool
|
||||||
|
|
||||||
//TODO: we should have many more options for TCP, an option for ICMP, and mimic the kernel a bit better
|
//TODO: we should have many more options for TCP, an option for ICMP, and mimic the kernel a bit better
|
||||||
// https://www.kernel.org/doc/Documentation/networking/nf_conntrack-sysctl.txt
|
// https://www.kernel.org/doc/Documentation/networking/nf_conntrack-sysctl.txt
|
||||||
@@ -58,9 +58,8 @@ type Firewall struct {
|
|||||||
routableNetworks *bart.Lite
|
routableNetworks *bart.Lite
|
||||||
|
|
||||||
// assignedNetworks is a list of vpn networks assigned to us in the certificate.
|
// assignedNetworks is a list of vpn networks assigned to us in the certificate.
|
||||||
assignedNetworks []netip.Prefix
|
assignedNetworks []netip.Prefix
|
||||||
// unsafeNetworks is the list of unsafe networks issued to us in the certificate
|
hasUnsafeNetworks bool
|
||||||
unsafeNetworks []netip.Prefix
|
|
||||||
|
|
||||||
rules string
|
rules string
|
||||||
rulesVersion uint16
|
rulesVersion uint16
|
||||||
@@ -69,7 +68,15 @@ type Firewall struct {
|
|||||||
incomingMetrics firewallMetrics
|
incomingMetrics firewallMetrics
|
||||||
outgoingMetrics firewallMetrics
|
outgoingMetrics firewallMetrics
|
||||||
|
|
||||||
l *slog.Logger
|
// reporter is the optional embedder-supplied event sink. Immutable for
|
||||||
|
// the lifetime of this Firewall; Control.SetFirewallEventReporter
|
||||||
|
// installs it by shallow-copying the Firewall under the conntrack lock
|
||||||
|
// and swapping the pointer, and reloadFirewall carries it forward.
|
||||||
|
// Read unsynchronized on the data path: the preceding Firewall-pointer
|
||||||
|
// read pins the field's value for the duration of that call.
|
||||||
|
reporter events.Reporter
|
||||||
|
|
||||||
|
l *logrus.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
type firewallMetrics struct {
|
type firewallMetrics struct {
|
||||||
@@ -133,7 +140,7 @@ type firewallLocalCIDR struct {
|
|||||||
|
|
||||||
// NewFirewall creates a new Firewall object. A TimerWheel is created for you from the provided timeouts.
|
// NewFirewall creates a new Firewall object. A TimerWheel is created for you from the provided timeouts.
|
||||||
// The certificate provided should be the highest version loaded in memory.
|
// The certificate provided should be the highest version loaded in memory.
|
||||||
func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Duration, c cert.Certificate) *Firewall {
|
func NewFirewall(l *logrus.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Duration, c cert.Certificate) *Firewall {
|
||||||
//TODO: error on 0 duration
|
//TODO: error on 0 duration
|
||||||
var tmin, tmax time.Duration
|
var tmin, tmax time.Duration
|
||||||
|
|
||||||
@@ -159,9 +166,10 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
|
|||||||
assignedNetworks = append(assignedNetworks, network)
|
assignedNetworks = append(assignedNetworks, network)
|
||||||
}
|
}
|
||||||
|
|
||||||
unsafeNetworks := c.UnsafeNetworks()
|
hasUnsafeNetworks := false
|
||||||
for _, n := range unsafeNetworks {
|
for _, n := range c.UnsafeNetworks() {
|
||||||
routableNetworks.Insert(n)
|
routableNetworks.Insert(n)
|
||||||
|
hasUnsafeNetworks = true
|
||||||
}
|
}
|
||||||
|
|
||||||
return &Firewall{
|
return &Firewall{
|
||||||
@@ -169,15 +177,15 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
|
|||||||
Conns: make(map[firewall.Packet]*conn),
|
Conns: make(map[firewall.Packet]*conn),
|
||||||
TimerWheel: NewTimerWheel[firewall.Packet](tmin, tmax),
|
TimerWheel: NewTimerWheel[firewall.Packet](tmin, tmax),
|
||||||
},
|
},
|
||||||
InRules: newFirewallTable(),
|
InRules: newFirewallTable(),
|
||||||
OutRules: newFirewallTable(),
|
OutRules: newFirewallTable(),
|
||||||
TCPTimeout: tcpTimeout,
|
TCPTimeout: tcpTimeout,
|
||||||
UDPTimeout: UDPTimeout,
|
UDPTimeout: UDPTimeout,
|
||||||
DefaultTimeout: defaultTimeout,
|
DefaultTimeout: defaultTimeout,
|
||||||
routableNetworks: routableNetworks,
|
routableNetworks: routableNetworks,
|
||||||
assignedNetworks: assignedNetworks,
|
assignedNetworks: assignedNetworks,
|
||||||
unsafeNetworks: unsafeNetworks,
|
hasUnsafeNetworks: hasUnsafeNetworks,
|
||||||
l: l,
|
l: l,
|
||||||
|
|
||||||
incomingMetrics: firewallMetrics{
|
incomingMetrics: firewallMetrics{
|
||||||
droppedLocalAddr: metrics.GetOrRegisterCounter("firewall.incoming.dropped.local_addr", nil),
|
droppedLocalAddr: metrics.GetOrRegisterCounter("firewall.incoming.dropped.local_addr", nil),
|
||||||
@@ -192,7 +200,7 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewFirewallFromConfig(l *slog.Logger, cs *CertState, c *config.C) (*Firewall, error) {
|
func NewFirewallFromConfig(l *logrus.Logger, cs *CertState, c *config.C) (*Firewall, error) {
|
||||||
certificate := cs.getCertificate(cert.Version2)
|
certificate := cs.getCertificate(cert.Version2)
|
||||||
if certificate == nil {
|
if certificate == nil {
|
||||||
certificate = cs.getCertificate(cert.Version1)
|
certificate = cs.getCertificate(cert.Version1)
|
||||||
@@ -216,23 +224,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.InboundSendReject = true
|
fw.InSendReject = true
|
||||||
case "drop":
|
case "drop":
|
||||||
fw.InboundSendReject = false
|
fw.InSendReject = false
|
||||||
default:
|
default:
|
||||||
l.Warn("invalid firewall.inbound_action, defaulting to `drop`", "action", inboundAction)
|
l.WithField("action", inboundAction).Warn("invalid firewall.inbound_action, defaulting to `drop`")
|
||||||
fw.InboundSendReject = false
|
fw.InSendReject = false
|
||||||
}
|
}
|
||||||
|
|
||||||
outboundAction := c.GetString("firewall.outbound_action", "drop")
|
outboundAction := c.GetString("firewall.outbound_action", "drop")
|
||||||
switch outboundAction {
|
switch outboundAction {
|
||||||
case "reject":
|
case "reject":
|
||||||
fw.OutboundSendReject = true
|
fw.OutSendReject = true
|
||||||
case "drop":
|
case "drop":
|
||||||
fw.OutboundSendReject = false
|
fw.OutSendReject = false
|
||||||
default:
|
default:
|
||||||
l.Warn("invalid firewall.outbound_action, defaulting to `drop`", "action", outboundAction)
|
l.WithField("action", outboundAction).Warn("invalid firewall.outbound_action, defaulting to `drop`")
|
||||||
fw.OutboundSendReject = false
|
fw.OutSendReject = false
|
||||||
}
|
}
|
||||||
|
|
||||||
err := AddFirewallRulesFromConfig(l, false, c, fw)
|
err := AddFirewallRulesFromConfig(l, false, c, fw)
|
||||||
@@ -269,7 +277,7 @@ func (f *Firewall) AddRule(incoming bool, proto uint8, startPort int32, endPort
|
|||||||
case firewall.ProtoICMP, firewall.ProtoICMPv6:
|
case firewall.ProtoICMP, firewall.ProtoICMPv6:
|
||||||
//ICMP traffic doesn't have ports, so we always coerce to "any", even if a value is provided
|
//ICMP traffic doesn't have ports, so we always coerce to "any", even if a value is provided
|
||||||
if startPort != firewall.PortAny {
|
if startPort != firewall.PortAny {
|
||||||
f.l.Warn("ignoring port specification for ICMP firewall rule", "startPort", startPort)
|
f.l.WithField("startPort", startPort).Warn("ignoring port specification for ICMP firewall rule")
|
||||||
}
|
}
|
||||||
startPort = firewall.PortAny
|
startPort = firewall.PortAny
|
||||||
endPort = firewall.PortAny
|
endPort = firewall.PortAny
|
||||||
@@ -291,9 +299,8 @@ func (f *Firewall) AddRule(incoming bool, proto uint8, startPort int32, endPort
|
|||||||
if !incoming {
|
if !incoming {
|
||||||
direction = "outgoing"
|
direction = "outgoing"
|
||||||
}
|
}
|
||||||
f.l.Info("Firewall rule added",
|
f.l.WithField("firewallRule", m{"direction": direction, "proto": proto, "startPort": startPort, "endPort": endPort, "groups": groups, "host": host, "cidr": cidr, "localCidr": localCidr, "caName": caName, "caSha": caSha}).
|
||||||
"firewallRule", m{"direction": direction, "proto": proto, "startPort": startPort, "endPort": endPort, "groups": groups, "host": host, "cidr": cidr, "localCidr": localCidr, "caName": caName, "caSha": caSha},
|
Info("Firewall rule added")
|
||||||
)
|
|
||||||
|
|
||||||
return fp.addRule(f, startPort, endPort, groups, host, cidr, localCidr, caName, caSha)
|
return fp.addRule(f, startPort, endPort, groups, host, cidr, localCidr, caName, caSha)
|
||||||
}
|
}
|
||||||
@@ -316,7 +323,7 @@ func (f *Firewall) GetRuleHashes() string {
|
|||||||
return "SHA:" + f.GetRuleHash() + ",FNV:" + strconv.FormatUint(uint64(f.GetRuleHashFNV()), 10)
|
return "SHA:" + f.GetRuleHash() + ",FNV:" + strconv.FormatUint(uint64(f.GetRuleHashFNV()), 10)
|
||||||
}
|
}
|
||||||
|
|
||||||
func AddFirewallRulesFromConfig(l *slog.Logger, inbound bool, c *config.C, fw FirewallInterface) error {
|
func AddFirewallRulesFromConfig(l *logrus.Logger, inbound bool, c *config.C, fw FirewallInterface) error {
|
||||||
var table string
|
var table string
|
||||||
if inbound {
|
if inbound {
|
||||||
table = "firewall.inbound"
|
table = "firewall.inbound"
|
||||||
@@ -374,7 +381,7 @@ func AddFirewallRulesFromConfig(l *slog.Logger, inbound bool, c *config.C, fw Fi
|
|||||||
startPort = firewall.PortAny
|
startPort = firewall.PortAny
|
||||||
endPort = firewall.PortAny
|
endPort = firewall.PortAny
|
||||||
if sPort != "" {
|
if sPort != "" {
|
||||||
l.Warn("ignoring port specification for ICMP firewall rule", "port", sPort)
|
l.WithField("port", sPort).Warn("ignoring port specification for ICMP firewall rule")
|
||||||
}
|
}
|
||||||
default:
|
default:
|
||||||
return fmt.Errorf("%s rule #%v; proto was not understood; `%s`", table, i, r.Proto)
|
return fmt.Errorf("%s rule #%v; proto was not understood; `%s`", table, i, r.Proto)
|
||||||
@@ -398,11 +405,7 @@ func AddFirewallRulesFromConfig(l *slog.Logger, inbound bool, c *config.C, fw Fi
|
|||||||
}
|
}
|
||||||
|
|
||||||
if warning := r.sanity(); warning != nil {
|
if warning := r.sanity(); warning != nil {
|
||||||
l.Warn("firewall rule sanity check",
|
l.Warnf("%s rule #%v; %s", table, i, warning)
|
||||||
"table", table,
|
|
||||||
"rule", i,
|
|
||||||
"warning", warning,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
err = fw.AddRule(inbound, proto, startPort, endPort, r.Groups, r.Host, r.Cidr, r.LocalCidr, r.CAName, r.CASha)
|
err = fw.AddRule(inbound, proto, startPort, endPort, r.Groups, r.Host, r.Cidr, r.LocalCidr, r.CAName, r.CASha)
|
||||||
@@ -422,18 +425,27 @@ 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, ctx firewall.PacketContext, 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
|
||||||
|
}
|
||||||
|
|
||||||
|
peerCert := h.ConnectionState.peerCert
|
||||||
|
|
||||||
// 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
|
||||||
if h.vpnAddrs[0] != fp.RemoteAddr {
|
if h.vpnAddrs[0] != fp.RemoteAddr {
|
||||||
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
||||||
|
f.reportDrop(incoming, events.DropInvalidRemoteIP, fp, ctx, peerCert)
|
||||||
return ErrInvalidRemoteIP
|
return ErrInvalidRemoteIP
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
nwType, ok := h.networks.Lookup(fp.RemoteAddr)
|
nwType, ok := h.networks.Lookup(fp.RemoteAddr)
|
||||||
if !ok {
|
if !ok {
|
||||||
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
||||||
|
f.reportDrop(incoming, events.DropInvalidRemoteIP, fp, ctx, peerCert)
|
||||||
return ErrInvalidRemoteIP
|
return ErrInvalidRemoteIP
|
||||||
}
|
}
|
||||||
switch nwType {
|
switch nwType {
|
||||||
@@ -441,11 +453,13 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *
|
|||||||
break // nothing special
|
break // nothing special
|
||||||
case NetworkTypeVPNPeer:
|
case NetworkTypeVPNPeer:
|
||||||
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
||||||
|
f.reportDrop(incoming, events.DropPeerRejected, fp, ctx, peerCert)
|
||||||
return ErrPeerRejected // reject for now, one day this may have different FW rules
|
return ErrPeerRejected // reject for now, one day this may have different FW rules
|
||||||
case NetworkTypeUnsafe:
|
case NetworkTypeUnsafe:
|
||||||
break // nothing special, one day this may have different FW rules
|
break // nothing special, one day this may have different FW rules
|
||||||
default:
|
default:
|
||||||
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
||||||
|
f.reportDrop(incoming, events.DropUnknownNetwork, fp, ctx, peerCert)
|
||||||
return ErrUnknownNetworkType //should never happen
|
return ErrUnknownNetworkType //should never happen
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -453,27 +467,24 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *
|
|||||||
// Make sure we are supposed to be handling this local ip address
|
// Make sure we are supposed to be handling this local ip address
|
||||||
if !f.routableNetworks.Contains(fp.LocalAddr) {
|
if !f.routableNetworks.Contains(fp.LocalAddr) {
|
||||||
f.metrics(incoming).droppedLocalAddr.Inc(1)
|
f.metrics(incoming).droppedLocalAddr.Inc(1)
|
||||||
|
f.reportDrop(incoming, events.DropInvalidLocalIP, fp, ctx, peerCert)
|
||||||
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
|
||||||
}
|
}
|
||||||
|
|
||||||
// We now know which firewall table to check against
|
// We now know which firewall table to check against
|
||||||
if !table.match(fp, incoming, h.ConnectionState.peerCert, caPool) {
|
if !table.match(fp, incoming, peerCert, caPool) {
|
||||||
f.metrics(incoming).droppedNoRule.Inc(1)
|
f.metrics(incoming).droppedNoRule.Inc(1)
|
||||||
|
f.reportDrop(incoming, events.DropNoMatchingRule, fp, ctx, peerCert)
|
||||||
return ErrNoMatchingRule
|
return ErrNoMatchingRule
|
||||||
}
|
}
|
||||||
|
|
||||||
// We always want to conntrack since it is a faster operation
|
// We always want to conntrack since it is a faster operation
|
||||||
f.addConn(fp, incoming)
|
f.addConn(fp, ctx, incoming, peerCert)
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -492,6 +503,59 @@ func (f *Firewall) Destroy() {
|
|||||||
//TODO: clean references if/when needed
|
//TODO: clean references if/when needed
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (f *Firewall) reportDrop(incoming bool, reason events.DropReason, fp firewall.Packet, ctx firewall.PacketContext, peerCert *cert.CachedCertificate) {
|
||||||
|
r := f.reporter
|
||||||
|
if r == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.ReportDrop(events.DropEvent{
|
||||||
|
Incoming: incoming,
|
||||||
|
Reason: reason,
|
||||||
|
Packet: fp,
|
||||||
|
Context: ctx,
|
||||||
|
PeerCert: peerCert,
|
||||||
|
RulesVersion: f.rulesVersion,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Firewall) reportFlowCreate(incoming bool, fp firewall.Packet, ctx firewall.PacketContext, peerCert *cert.CachedCertificate) {
|
||||||
|
r := f.reporter
|
||||||
|
if r == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.ReportFlowCreate(events.FlowCreateEvent{
|
||||||
|
Incoming: incoming,
|
||||||
|
Packet: fp,
|
||||||
|
Context: ctx,
|
||||||
|
PeerCert: peerCert,
|
||||||
|
RulesVersion: f.rulesVersion,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Firewall) reportFlowEvict(incoming bool, fp firewall.Packet, rulesVersion uint16, expired bool) {
|
||||||
|
r := f.reporter
|
||||||
|
if r == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.ReportFlowEvict(events.FlowEvictEvent{
|
||||||
|
Incoming: incoming,
|
||||||
|
Packet: fp,
|
||||||
|
RulesVersion: rulesVersion,
|
||||||
|
Expired: expired,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Firewall) reportRulesReload(oldVersion, newVersion uint16) {
|
||||||
|
r := f.reporter
|
||||||
|
if r == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.ReportRulesReload(events.RulesReloadEvent{
|
||||||
|
OldVersion: oldVersion,
|
||||||
|
NewVersion: newVersion,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Firewall) EmitStats() {
|
func (f *Firewall) EmitStats() {
|
||||||
conntrack := f.Conntrack
|
conntrack := f.Conntrack
|
||||||
conntrack.Lock()
|
conntrack.Lock()
|
||||||
@@ -534,26 +598,28 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
|
|||||||
|
|
||||||
// We now know which firewall table to check against
|
// We now know which firewall table to check against
|
||||||
if !table.match(fp, c.incoming, h.ConnectionState.peerCert, caPool) {
|
if !table.match(fp, c.incoming, h.ConnectionState.peerCert, caPool) {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Level >= logrus.DebugLevel {
|
||||||
h.logger(f.l).Debug("dropping old conntrack entry, does not match new ruleset",
|
h.logger(f.l).
|
||||||
"fwPacket", fp,
|
WithField("fwPacket", fp).
|
||||||
"incoming", c.incoming,
|
WithField("incoming", c.incoming).
|
||||||
"rulesVersion", f.rulesVersion,
|
WithField("rulesVersion", f.rulesVersion).
|
||||||
"oldRulesVersion", c.rulesVersion,
|
WithField("oldRulesVersion", c.rulesVersion).
|
||||||
)
|
Debugln("dropping old conntrack entry, does not match new ruleset")
|
||||||
}
|
}
|
||||||
|
oldRulesVersion := c.rulesVersion
|
||||||
delete(conntrack.Conns, fp)
|
delete(conntrack.Conns, fp)
|
||||||
|
f.reportFlowEvict(c.incoming, fp, oldRulesVersion, false)
|
||||||
conntrack.Unlock()
|
conntrack.Unlock()
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Level >= logrus.DebugLevel {
|
||||||
h.logger(f.l).Debug("keeping old conntrack entry, does match new ruleset",
|
h.logger(f.l).
|
||||||
"fwPacket", fp,
|
WithField("fwPacket", fp).
|
||||||
"incoming", c.incoming,
|
WithField("incoming", c.incoming).
|
||||||
"rulesVersion", f.rulesVersion,
|
WithField("rulesVersion", f.rulesVersion).
|
||||||
"oldRulesVersion", c.rulesVersion,
|
WithField("oldRulesVersion", c.rulesVersion).
|
||||||
)
|
Debugln("keeping old conntrack entry, does match new ruleset")
|
||||||
}
|
}
|
||||||
|
|
||||||
c.rulesVersion = f.rulesVersion
|
c.rulesVersion = f.rulesVersion
|
||||||
@@ -577,7 +643,7 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
|
func (f *Firewall) addConn(fp firewall.Packet, ctx firewall.PacketContext, incoming bool, peerCert *cert.CachedCertificate) {
|
||||||
var timeout time.Duration
|
var timeout time.Duration
|
||||||
c := &conn{}
|
c := &conn{}
|
||||||
|
|
||||||
@@ -592,7 +658,8 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
|
|||||||
|
|
||||||
conntrack := f.Conntrack
|
conntrack := f.Conntrack
|
||||||
conntrack.Lock()
|
conntrack.Lock()
|
||||||
if _, ok := conntrack.Conns[fp]; !ok {
|
_, existing := conntrack.Conns[fp]
|
||||||
|
if !existing {
|
||||||
conntrack.TimerWheel.Advance(time.Now())
|
conntrack.TimerWheel.Advance(time.Now())
|
||||||
conntrack.TimerWheel.Add(fp, timeout)
|
conntrack.TimerWheel.Add(fp, timeout)
|
||||||
}
|
}
|
||||||
@@ -603,6 +670,13 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
|
|||||||
c.rulesVersion = f.rulesVersion
|
c.rulesVersion = f.rulesVersion
|
||||||
c.Expires = time.Now().Add(timeout)
|
c.Expires = time.Now().Add(timeout)
|
||||||
conntrack.Conns[fp] = c
|
conntrack.Conns[fp] = c
|
||||||
|
|
||||||
|
// Report only when this represents a genuinely new flow. Fires under the
|
||||||
|
// conntrack lock so FlowCreate/FlowEvict events stay ordered relative to
|
||||||
|
// RulesReloadEvent, which also fires under this lock.
|
||||||
|
if !existing {
|
||||||
|
f.reportFlowCreate(incoming, fp, ctx, peerCert)
|
||||||
|
}
|
||||||
conntrack.Unlock()
|
conntrack.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -626,7 +700,10 @@ func (f *Firewall) evict(p firewall.Packet) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// This conn is done
|
// This conn is done
|
||||||
|
rulesVersion := t.rulesVersion
|
||||||
|
incoming := t.incoming
|
||||||
delete(conntrack.Conns, p)
|
delete(conntrack.Conns, p)
|
||||||
|
f.reportFlowEvict(incoming, p, rulesVersion, true)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedCertificate, caPool *cert.CAPool) bool {
|
func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedCertificate, caPool *cert.CAPool) bool {
|
||||||
@@ -897,7 +974,7 @@ func (flc *firewallLocalCIDR) addRule(f *Firewall, localCidr string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if localCidr == "" {
|
if localCidr == "" {
|
||||||
if len(f.unsafeNetworks) == 0 || f.defaultLocalCIDRAny {
|
if !f.hasUnsafeNetworks || f.defaultLocalCIDRAny {
|
||||||
flc.Any = true
|
flc.Any = true
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -941,7 +1018,7 @@ type rule struct {
|
|||||||
CASha string
|
CASha string
|
||||||
}
|
}
|
||||||
|
|
||||||
func convertRule(l *slog.Logger, p any, table string, i int) (rule, error) {
|
func convertRule(l *logrus.Logger, p any, table string, i int) (rule, error) {
|
||||||
r := rule{}
|
r := rule{}
|
||||||
|
|
||||||
m, ok := p.(map[string]any)
|
m, ok := p.(map[string]any)
|
||||||
@@ -972,10 +1049,7 @@ func convertRule(l *slog.Logger, p any, table string, i int) (rule, error) {
|
|||||||
return r, errors.New("group should contain a single value, an array with more than one entry was provided")
|
return r, errors.New("group should contain a single value, an array with more than one entry was provided")
|
||||||
}
|
}
|
||||||
|
|
||||||
l.Warn("group was an array with a single value, converting to simple value",
|
l.Warnf("%s rule #%v; group was an array with a single value, converting to simple value", table, i)
|
||||||
"table", table,
|
|
||||||
"rule", i,
|
|
||||||
)
|
|
||||||
m["group"] = v[0]
|
m["group"] = v[0]
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1055,6 +1129,7 @@ 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
|
||||||
@@ -1063,11 +1138,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 := parsePortValue("", s)
|
rPort, err := strconv.Atoi(s)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return notAPort, notAPort, err
|
return notAPort, notAPort, fmt.Errorf("was not a number; `%s`", s)
|
||||||
}
|
}
|
||||||
return rPort, rPort, nil
|
return int32(rPort), int32(rPort), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
sPorts := strings.SplitN(s, `-`, 2)
|
sPorts := strings.SplitN(s, `-`, 2)
|
||||||
@@ -1078,40 +1153,22 @@ 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)
|
||||||
}
|
}
|
||||||
|
|
||||||
startPort, err := parsePortValue("beginning range ", sPorts[0])
|
rStartPort, err := strconv.Atoi(sPorts[0])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return notAPort, notAPort, err
|
return notAPort, notAPort, fmt.Errorf("beginning range was not a number; `%s`", sPorts[0])
|
||||||
}
|
}
|
||||||
|
|
||||||
endPort, err := parsePortValue("ending range ", sPorts[1])
|
rEndPort, err := strconv.Atoi(sPorts[1])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return notAPort, notAPort, err
|
return notAPort, notAPort, fmt.Errorf("ending range was not a number; `%s`", sPorts[1])
|
||||||
}
|
}
|
||||||
|
|
||||||
|
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)
|
|
||||||
}
|
|
||||||
|
|||||||
+6
-7
@@ -2,9 +2,10 @@ package firewall
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"log/slog"
|
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ConntrackCache is used as a local routine cache to know if a given flow
|
// ConntrackCache is used as a local routine cache to know if a given flow
|
||||||
@@ -15,17 +16,15 @@ type ConntrackCacheTicker struct {
|
|||||||
cacheV uint64
|
cacheV uint64
|
||||||
cacheTick atomic.Uint64
|
cacheTick atomic.Uint64
|
||||||
|
|
||||||
l *slog.Logger
|
|
||||||
cache ConntrackCache
|
cache ConntrackCache
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewConntrackCacheTicker(ctx context.Context, l *slog.Logger, d time.Duration) *ConntrackCacheTicker {
|
func NewConntrackCacheTicker(ctx context.Context, d time.Duration) *ConntrackCacheTicker {
|
||||||
if d == 0 {
|
if d == 0 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
c := &ConntrackCacheTicker{
|
c := &ConntrackCacheTicker{
|
||||||
l: l,
|
|
||||||
cache: ConntrackCache{},
|
cache: ConntrackCache{},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -49,15 +48,15 @@ func (c *ConntrackCacheTicker) tick(ctx context.Context, d time.Duration) {
|
|||||||
|
|
||||||
// Get checks if the cache ticker has moved to the next version before returning
|
// Get checks if the cache ticker has moved to the next version before returning
|
||||||
// the map. If it has moved, we reset the map.
|
// the map. If it has moved, we reset the map.
|
||||||
func (c *ConntrackCacheTicker) Get() ConntrackCache {
|
func (c *ConntrackCacheTicker) Get(l *logrus.Logger) ConntrackCache {
|
||||||
if c == nil {
|
if c == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
||||||
c.cacheV = tick
|
c.cacheV = tick
|
||||||
if ll := len(c.cache); ll > 0 {
|
if ll := len(c.cache); ll > 0 {
|
||||||
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
if l.Level == logrus.DebugLevel {
|
||||||
c.l.Debug("resetting conntrack cache", "len", ll)
|
l.WithField("len", ll).Debug("resetting conntrack cache")
|
||||||
}
|
}
|
||||||
c.cache = make(ConntrackCache, ll)
|
c.cache = make(ConntrackCache, ll)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,69 +0,0 @@
|
|||||||
package firewall
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"log/slog"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
// The tests below pin the log format produced by ConntrackCacheTicker.Get
|
|
||||||
// so changes cannot silently break what operators are grepping for. The
|
|
||||||
// ticker's internal state (cache + cacheTick) is poked directly to avoid
|
|
||||||
// racing a goroutine-driven tick in tests.
|
|
||||||
|
|
||||||
func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheTicker {
|
|
||||||
t.Helper()
|
|
||||||
c := &ConntrackCacheTicker{
|
|
||||||
l: l,
|
|
||||||
cache: make(ConntrackCache, cacheLen),
|
|
||||||
}
|
|
||||||
for i := 0; i < cacheLen; i++ {
|
|
||||||
c.cache[Packet{LocalPort: uint16(i) + 1}] = struct{}{}
|
|
||||||
}
|
|
||||||
c.cacheTick.Store(1) // cacheV starts at 0, so Get() takes the reset path
|
|
||||||
return c
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
|
||||||
buf := &bytes.Buffer{}
|
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 3)
|
|
||||||
c.Get()
|
|
||||||
|
|
||||||
assert.Equal(t, "level=DEBUG msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
|
||||||
buf := &bytes.Buffer{}
|
|
||||||
l := test.NewJSONLoggerWithOutput(buf, slog.LevelDebug)
|
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 2)
|
|
||||||
c.Get()
|
|
||||||
|
|
||||||
assert.JSONEq(t, `{"level":"DEBUG","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
|
||||||
buf := &bytes.Buffer{}
|
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelInfo)
|
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 5)
|
|
||||||
c.Get()
|
|
||||||
|
|
||||||
assert.Empty(t, buf.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
|
||||||
buf := &bytes.Buffer{}
|
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 0)
|
|
||||||
c.Get()
|
|
||||||
|
|
||||||
assert.Empty(t, buf.String())
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,102 @@
|
|||||||
|
// Package events defines the opt-in firewall event reporting interface.
|
||||||
|
//
|
||||||
|
// Nebula emits raw packet-level events (drops, flow creations, flow evictions,
|
||||||
|
// rule reloads) and does no aggregation, counting, batching, rule-description,
|
||||||
|
// transport, or timestamping. Embedders correlate events back to yaml rules
|
||||||
|
// out of band and capture whatever clock they need themselves. All Report*
|
||||||
|
// methods are invoked while nebula holds internal locks and must be
|
||||||
|
// non-blocking.
|
||||||
|
//
|
||||||
|
// Events are passed to Report* methods by value. Implementations must not
|
||||||
|
// take the address of a received event: doing so forces Go's escape
|
||||||
|
// analysis to move the event to the heap and costs one allocation per call.
|
||||||
|
// To forward an event, either copy its fields into the reporter's own
|
||||||
|
// pooled record or send it through a value-typed channel (chan DropEvent,
|
||||||
|
// not chan *DropEvent).
|
||||||
|
package events
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/firewall"
|
||||||
|
)
|
||||||
|
|
||||||
|
type DropReason uint8
|
||||||
|
|
||||||
|
const (
|
||||||
|
DropInvalidLocalIP DropReason = iota
|
||||||
|
DropInvalidRemoteIP
|
||||||
|
DropPeerRejected
|
||||||
|
DropUnknownNetwork
|
||||||
|
DropNoMatchingRule
|
||||||
|
)
|
||||||
|
|
||||||
|
func (r DropReason) String() string {
|
||||||
|
switch r {
|
||||||
|
case DropInvalidLocalIP:
|
||||||
|
return "invalid_local_ip"
|
||||||
|
case DropInvalidRemoteIP:
|
||||||
|
return "invalid_remote_ip"
|
||||||
|
case DropPeerRejected:
|
||||||
|
return "peer_rejected"
|
||||||
|
case DropUnknownNetwork:
|
||||||
|
return "unknown_network"
|
||||||
|
case DropNoMatchingRule:
|
||||||
|
return "no_matching_rule"
|
||||||
|
default:
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// DropEvent is emitted for every packet that fails the firewall check. Drops
|
||||||
|
// are not aggregated; every drop produces one event.
|
||||||
|
type DropEvent struct {
|
||||||
|
Incoming bool
|
||||||
|
Reason DropReason
|
||||||
|
Packet firewall.Packet
|
||||||
|
Context firewall.PacketContext
|
||||||
|
PeerCert *cert.CachedCertificate
|
||||||
|
RulesVersion uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
// FlowCreateEvent is emitted when a packet is allowed and a new conntrack
|
||||||
|
// entry is created. Subsequent packets in the same flow do not re-emit.
|
||||||
|
type FlowCreateEvent struct {
|
||||||
|
Incoming bool
|
||||||
|
Packet firewall.Packet
|
||||||
|
Context firewall.PacketContext
|
||||||
|
PeerCert *cert.CachedCertificate
|
||||||
|
RulesVersion uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
// FlowEvictEvent is emitted when a conntrack entry is removed. Context is
|
||||||
|
// not carried: timer-wheel eviction has no packet in hand, and reload
|
||||||
|
// revalidation evicts the OLD flow rather than the triggering packet.
|
||||||
|
// RulesVersion is the version under which the flow was originally allowed,
|
||||||
|
// which may differ from the current firewall version.
|
||||||
|
type FlowEvictEvent struct {
|
||||||
|
Incoming bool
|
||||||
|
Packet firewall.Packet
|
||||||
|
RulesVersion uint16
|
||||||
|
// Expired is true when eviction was due to conntrack timeout; false when
|
||||||
|
// the entry was removed because it failed re-validation after a reload.
|
||||||
|
Expired bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// RulesReloadEvent is emitted once after each successful firewall reload.
|
||||||
|
// Reporters that bucket state by RulesVersion should close the old bucket
|
||||||
|
// and open a new one on receipt.
|
||||||
|
type RulesReloadEvent struct {
|
||||||
|
OldVersion uint16
|
||||||
|
NewVersion uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reporter is the embedder-supplied sink for firewall events. Implementations
|
||||||
|
// that want a timestamp should call time.Now() themselves at the top of the
|
||||||
|
// method; nebula does not provide one. See the package doc for the
|
||||||
|
// do-not-take-address rule.
|
||||||
|
type Reporter interface {
|
||||||
|
ReportDrop(DropEvent)
|
||||||
|
ReportFlowCreate(FlowCreateEvent)
|
||||||
|
ReportFlowEvict(FlowEvictEvent)
|
||||||
|
ReportRulesReload(RulesReloadEvent)
|
||||||
|
}
|
||||||
@@ -31,6 +31,27 @@ type Packet struct {
|
|||||||
Fragment bool
|
Fragment bool
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// PacketContext carries additional parsed details about a packet that are
|
||||||
|
// useful for event reporting but deliberately kept out of Packet so Packet
|
||||||
|
// can keep being used as a conntrack map key. Populated alongside Packet by
|
||||||
|
// newPacket.
|
||||||
|
//
|
||||||
|
// Fields are interpreted based on Packet.Protocol:
|
||||||
|
// - ProtoTCP: TCPFlags is meaningful; ICMPType / ICMPCode are zero
|
||||||
|
// - ProtoICMP, ProtoICMPv6: ICMPType / ICMPCode are meaningful; TCPFlags is zero
|
||||||
|
// - ProtoUDP and others: only Length is meaningful
|
||||||
|
type PacketContext struct {
|
||||||
|
// Length is the total IP packet length in bytes, including headers.
|
||||||
|
Length uint16
|
||||||
|
// TCPFlags is the flag byte from the TCP header (bits for FIN, SYN, RST,
|
||||||
|
// PSH, ACK, URG, ECE, CWR).
|
||||||
|
TCPFlags uint8
|
||||||
|
// ICMPType is the type field of the ICMP / ICMPv6 header.
|
||||||
|
ICMPType uint8
|
||||||
|
// ICMPCode is the code field of the ICMP / ICMPv6 header.
|
||||||
|
ICMPCode uint8
|
||||||
|
}
|
||||||
|
|
||||||
func (fp *Packet) Copy() *Packet {
|
func (fp *Packet) Copy() *Packet {
|
||||||
return &Packet{
|
return &Packet{
|
||||||
LocalAddr: fp.LocalAddr,
|
LocalAddr: fp.LocalAddr,
|
||||||
|
|||||||
@@ -0,0 +1,731 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gaissmai/bart"
|
||||||
|
"github.com/google/gopacket"
|
||||||
|
"github.com/google/gopacket/layers"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/firewall"
|
||||||
|
"github.com/slackhq/nebula/firewall/events"
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// recordingReporter captures every event fired against it. Its methods take
|
||||||
|
// the conntrack lock implicitly (via the firewall code path that invokes
|
||||||
|
// them), so we synchronize accumulator mutations with a small mutex to keep
|
||||||
|
// the race detector happy across goroutines in case a test introduces any.
|
||||||
|
type recordingReporter struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
drops []recordedDrop
|
||||||
|
creates []recordedCreate
|
||||||
|
evicts []recordedEvict
|
||||||
|
reloads []recordedReload
|
||||||
|
}
|
||||||
|
|
||||||
|
type recordedDrop struct {
|
||||||
|
incoming bool
|
||||||
|
reason events.DropReason
|
||||||
|
remote netip.Addr
|
||||||
|
local netip.Addr
|
||||||
|
peerName string
|
||||||
|
rulesVersion uint16
|
||||||
|
ctx firewall.PacketContext
|
||||||
|
}
|
||||||
|
|
||||||
|
type recordedCreate struct {
|
||||||
|
incoming bool
|
||||||
|
remote netip.Addr
|
||||||
|
local netip.Addr
|
||||||
|
peerName string
|
||||||
|
rulesVersion uint16
|
||||||
|
ctx firewall.PacketContext
|
||||||
|
}
|
||||||
|
|
||||||
|
type recordedEvict struct {
|
||||||
|
incoming bool
|
||||||
|
remote netip.Addr
|
||||||
|
local netip.Addr
|
||||||
|
rulesVersion uint16
|
||||||
|
expired bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type recordedReload struct {
|
||||||
|
oldVersion uint16
|
||||||
|
newVersion uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *recordingReporter) ReportDrop(e events.DropEvent) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
name := ""
|
||||||
|
if e.PeerCert != nil && e.PeerCert.Certificate != nil {
|
||||||
|
name = e.PeerCert.Certificate.Name()
|
||||||
|
}
|
||||||
|
r.drops = append(r.drops, recordedDrop{
|
||||||
|
incoming: e.Incoming,
|
||||||
|
reason: e.Reason,
|
||||||
|
remote: e.Packet.RemoteAddr,
|
||||||
|
local: e.Packet.LocalAddr,
|
||||||
|
peerName: name,
|
||||||
|
rulesVersion: e.RulesVersion,
|
||||||
|
ctx: e.Context,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *recordingReporter) ReportFlowCreate(e events.FlowCreateEvent) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
name := ""
|
||||||
|
if e.PeerCert != nil && e.PeerCert.Certificate != nil {
|
||||||
|
name = e.PeerCert.Certificate.Name()
|
||||||
|
}
|
||||||
|
r.creates = append(r.creates, recordedCreate{
|
||||||
|
incoming: e.Incoming,
|
||||||
|
remote: e.Packet.RemoteAddr,
|
||||||
|
local: e.Packet.LocalAddr,
|
||||||
|
peerName: name,
|
||||||
|
rulesVersion: e.RulesVersion,
|
||||||
|
ctx: e.Context,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *recordingReporter) ReportFlowEvict(e events.FlowEvictEvent) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.evicts = append(r.evicts, recordedEvict{
|
||||||
|
incoming: e.Incoming,
|
||||||
|
remote: e.Packet.RemoteAddr,
|
||||||
|
local: e.Packet.LocalAddr,
|
||||||
|
rulesVersion: e.RulesVersion,
|
||||||
|
expired: e.Expired,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *recordingReporter) ReportRulesReload(e events.RulesReloadEvent) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.reloads = append(r.reloads, recordedReload{
|
||||||
|
oldVersion: e.OldVersion,
|
||||||
|
newVersion: e.NewVersion,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// eventFixture builds a Firewall wired to a Control plus a packet/hostinfo
|
||||||
|
// pair that a test can reuse. By default the ruleset allows the packet;
|
||||||
|
// callers mutate fw / p / h as needed before invoking Drop.
|
||||||
|
type eventFixture struct {
|
||||||
|
ctl *Control
|
||||||
|
fw *Firewall
|
||||||
|
p firewall.Packet
|
||||||
|
h *HostInfo
|
||||||
|
cp *cert.CAPool
|
||||||
|
}
|
||||||
|
|
||||||
|
func newEventFixture(t *testing.T) *eventFixture {
|
||||||
|
t.Helper()
|
||||||
|
l := test.NewLogger()
|
||||||
|
|
||||||
|
// myVpnNetworksTable covers our single peer address so buildNetworks takes
|
||||||
|
// the "simple case" path (h.networks stays nil); tests that want a populated
|
||||||
|
// BART table overwrite h.networks directly.
|
||||||
|
vpnNetworks := new(bart.Lite)
|
||||||
|
vpnNetworks.Insert(netip.MustParsePrefix("1.2.3.0/24"))
|
||||||
|
|
||||||
|
// Use the same cert for "peer" and "local" endpoints, matching the
|
||||||
|
// TestFirewall_Drop fixture style: LocalAddr == RemoteAddr == peer vpn addr.
|
||||||
|
c := &dummyCert{
|
||||||
|
name: "host1",
|
||||||
|
networks: []netip.Prefix{netip.MustParsePrefix("1.2.3.4/24")},
|
||||||
|
groups: []string{"default-group"},
|
||||||
|
issuer: "signer-shasum",
|
||||||
|
}
|
||||||
|
h := &HostInfo{
|
||||||
|
ConnectionState: &ConnectionState{
|
||||||
|
peerCert: &cert.CachedCertificate{
|
||||||
|
Certificate: c,
|
||||||
|
InvertedGroups: map[string]struct{}{"default-group": {}},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
vpnAddrs: []netip.Addr{netip.MustParseAddr("1.2.3.4")},
|
||||||
|
}
|
||||||
|
h.buildNetworks(vpnNetworks, c)
|
||||||
|
|
||||||
|
fw := NewFirewall(l, time.Minute, time.Minute, time.Minute, c)
|
||||||
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", ""))
|
||||||
|
require.NoError(t, fw.AddRule(false, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", ""))
|
||||||
|
|
||||||
|
ctl := &Control{
|
||||||
|
f: &Interface{firewall: fw},
|
||||||
|
l: l,
|
||||||
|
}
|
||||||
|
|
||||||
|
return &eventFixture{
|
||||||
|
ctl: ctl,
|
||||||
|
fw: fw,
|
||||||
|
p: firewall.Packet{
|
||||||
|
LocalAddr: netip.MustParseAddr("1.2.3.4"),
|
||||||
|
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||||
|
LocalPort: 10,
|
||||||
|
RemotePort: 90,
|
||||||
|
Protocol: firewall.ProtoUDP,
|
||||||
|
},
|
||||||
|
h: h,
|
||||||
|
cp: cert.NewCAPool(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// firewall() returns the currently-installed firewall. Needed because
|
||||||
|
// SetFirewallEventReporter replaces it via shallow-copy swap.
|
||||||
|
func (f *eventFixture) firewall() *Firewall {
|
||||||
|
return f.ctl.f.firewall
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEvents_ReportDrop_InvalidRemoteIP(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
|
||||||
|
// Packet to an address not in the cert's networks.
|
||||||
|
f.p.RemoteAddr = netip.MustParseAddr("9.9.9.9")
|
||||||
|
assert.Equal(t, ErrInvalidRemoteIP, f.firewall().Drop(f.p, firewall.PacketContext{}, false, f.h, f.cp, nil))
|
||||||
|
|
||||||
|
require.Len(t, r.drops, 1)
|
||||||
|
assert.Equal(t, events.DropInvalidRemoteIP, r.drops[0].reason)
|
||||||
|
assert.False(t, r.drops[0].incoming)
|
||||||
|
assert.Equal(t, "host1", r.drops[0].peerName)
|
||||||
|
assert.Empty(t, r.creates)
|
||||||
|
assert.Empty(t, r.evicts)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEvents_ReportDrop_InvalidLocalIP(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
|
||||||
|
// LocalAddr outside our routable networks.
|
||||||
|
f.p.LocalAddr = netip.MustParseAddr("9.9.9.9")
|
||||||
|
assert.Equal(t, ErrInvalidLocalIP, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
|
||||||
|
require.Len(t, r.drops, 1)
|
||||||
|
assert.Equal(t, events.DropInvalidLocalIP, r.drops[0].reason)
|
||||||
|
assert.True(t, r.drops[0].incoming)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEvents_ReportDrop_NoMatchingRule(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
// Reset to a firewall with no matching rule.
|
||||||
|
l := test.NewLogger()
|
||||||
|
fw := NewFirewall(l, time.Minute, time.Minute, time.Minute, f.h.ConnectionState.peerCert.Certificate)
|
||||||
|
// Rule that won't match (group not in peer's groups).
|
||||||
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", ""))
|
||||||
|
require.NoError(t, fw.AddRule(false, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", ""))
|
||||||
|
f.ctl.f.firewall = fw
|
||||||
|
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
|
||||||
|
assert.Equal(t, ErrNoMatchingRule, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
require.Len(t, r.drops, 1)
|
||||||
|
assert.Equal(t, events.DropNoMatchingRule, r.drops[0].reason)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEvents_ReportDrop_PeerRejected(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
// Re-classify the remote as VPNPeer so it triggers DropPeerRejected.
|
||||||
|
f.h.networks = new(bart.Table[NetworkType])
|
||||||
|
f.h.networks.Insert(netip.MustParsePrefix("1.2.3.0/24"), NetworkTypeVPNPeer)
|
||||||
|
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
|
||||||
|
assert.Equal(t, ErrPeerRejected, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
require.Len(t, r.drops, 1)
|
||||||
|
assert.Equal(t, events.DropPeerRejected, r.drops[0].reason)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEvents_ReportDrop_UnknownNetwork(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
// Insert an unrecognized NetworkType value to hit the default branch.
|
||||||
|
f.h.networks = new(bart.Table[NetworkType])
|
||||||
|
f.h.networks.Insert(netip.MustParsePrefix("1.2.3.0/24"), NetworkTypeUnknown)
|
||||||
|
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
|
||||||
|
assert.Equal(t, ErrUnknownNetworkType, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
require.Len(t, r.drops, 1)
|
||||||
|
assert.Equal(t, events.DropUnknownNetwork, r.drops[0].reason)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEvents_ReportFlowCreate_OnceOnly(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
|
||||||
|
// First allowed packet creates the conntrack entry.
|
||||||
|
require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
// Second matching packet on the same tuple is short-circuited by conntrack
|
||||||
|
// and must not fire another FlowCreate.
|
||||||
|
require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
|
||||||
|
require.Len(t, r.creates, 1)
|
||||||
|
assert.True(t, r.creates[0].incoming)
|
||||||
|
assert.Equal(t, f.p.RemoteAddr, r.creates[0].remote)
|
||||||
|
assert.Empty(t, r.drops)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEvents_ReportFlowEvict_OnReloadPurge(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
|
||||||
|
// Create a flow under the current rules.
|
||||||
|
require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
require.Len(t, r.creates, 1)
|
||||||
|
|
||||||
|
// Simulate a reload that produces rules the existing flow no longer
|
||||||
|
// matches. Bump rulesVersion and replace InRules with an empty table so
|
||||||
|
// revalidation fails.
|
||||||
|
fw := f.firewall()
|
||||||
|
fw.Conntrack.Lock()
|
||||||
|
fw.rulesVersion++
|
||||||
|
fw.InRules = newFirewallTable()
|
||||||
|
fw.Conntrack.Unlock()
|
||||||
|
|
||||||
|
// Next packet triggers re-validation, which fails and evicts the entry.
|
||||||
|
err := fw.Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)
|
||||||
|
assert.Equal(t, ErrNoMatchingRule, err)
|
||||||
|
|
||||||
|
require.Len(t, r.evicts, 1)
|
||||||
|
assert.False(t, r.evicts[0].expired, "evict from reload purge is not expiration")
|
||||||
|
assert.True(t, r.evicts[0].incoming)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEvents_ReportFlowEvict_OnTimeout(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
|
||||||
|
require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
require.Len(t, r.creates, 1)
|
||||||
|
|
||||||
|
// Force expiration by rewinding the entry's deadline.
|
||||||
|
fw := f.firewall()
|
||||||
|
fw.Conntrack.Lock()
|
||||||
|
c := fw.Conntrack.Conns[f.p]
|
||||||
|
require.NotNil(t, c)
|
||||||
|
c.Expires = time.Now().Add(-time.Hour)
|
||||||
|
fw.evict(f.p)
|
||||||
|
fw.Conntrack.Unlock()
|
||||||
|
|
||||||
|
require.Len(t, r.evicts, 1)
|
||||||
|
assert.True(t, r.evicts[0].expired)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEvents_SetNil_Clears(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
require.Len(t, r.creates, 1)
|
||||||
|
|
||||||
|
f.ctl.SetFirewallEventReporter(nil)
|
||||||
|
resetConntrack(f.firewall())
|
||||||
|
require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
// No second create should be recorded.
|
||||||
|
assert.Len(t, r.creates, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEvents_ReporterSurvivesSwap(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
|
||||||
|
// Simulate a reload by swapping in a fresh Firewall that carries the
|
||||||
|
// reporter forward. Mirrors what reloadFirewall does with the shared
|
||||||
|
// conntrack pointer.
|
||||||
|
l := test.NewLogger()
|
||||||
|
oldFw := f.firewall()
|
||||||
|
newFw := NewFirewall(l, time.Minute, time.Minute, time.Minute, f.h.ConnectionState.peerCert.Certificate)
|
||||||
|
require.NoError(t, newFw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", ""))
|
||||||
|
require.NoError(t, newFw.AddRule(false, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", ""))
|
||||||
|
newFw.Conntrack = oldFw.Conntrack
|
||||||
|
newFw.rulesVersion = oldFw.rulesVersion + 1
|
||||||
|
newFw.reporter = oldFw.reporter
|
||||||
|
f.ctl.f.firewall = newFw
|
||||||
|
newFw.reportRulesReload(oldFw.rulesVersion, newFw.rulesVersion)
|
||||||
|
|
||||||
|
require.Len(t, r.reloads, 1)
|
||||||
|
assert.Equal(t, oldFw.rulesVersion, r.reloads[0].oldVersion)
|
||||||
|
assert.Equal(t, newFw.rulesVersion, r.reloads[0].newVersion)
|
||||||
|
|
||||||
|
// Events on the new firewall should still reach the same reporter.
|
||||||
|
require.NoError(t, newFw.Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
require.Len(t, r.creates, 1)
|
||||||
|
assert.Equal(t, newFw.rulesVersion, r.creates[0].rulesVersion)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEvents_InstallDoesNotMutateOldFirewall(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
before := f.firewall()
|
||||||
|
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
|
||||||
|
after := f.firewall()
|
||||||
|
assert.NotSame(t, before, after, "SetFirewallEventReporter must replace the Firewall pointer")
|
||||||
|
assert.Nil(t, before.reporter, "the pre-install Firewall must remain untouched")
|
||||||
|
assert.NotNil(t, after.reporter)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- PacketContext parse tests --------------------------------------------
|
||||||
|
|
||||||
|
func mustSerialize(t *testing.T, lrs ...gopacket.SerializableLayer) []byte {
|
||||||
|
t.Helper()
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
opt := gopacket.SerializeOptions{ComputeChecksums: false, FixLengths: true}
|
||||||
|
require.NoError(t, gopacket.SerializeLayers(buf, opt, lrs...))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPacketContext_IPv4_TCPFlags(t *testing.T) {
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4, TTL: 64, Protocol: layers.IPProtocolTCP,
|
||||||
|
SrcIP: net.IPv4(10, 0, 0, 1), DstIP: net.IPv4(10, 0, 0, 2),
|
||||||
|
}
|
||||||
|
tcp := &layers.TCP{SrcPort: 1234, DstPort: 80, SYN: true, ACK: true}
|
||||||
|
require.NoError(t, tcp.SetNetworkLayerForChecksum(ip))
|
||||||
|
data := mustSerialize(t, ip, tcp, gopacket.Payload([]byte("hello")))
|
||||||
|
|
||||||
|
var fp firewall.Packet
|
||||||
|
var ctx firewall.PacketContext
|
||||||
|
require.NoError(t, newPacket(data, true, &fp, &ctx))
|
||||||
|
|
||||||
|
assert.Equal(t, uint8(firewall.ProtoTCP), fp.Protocol)
|
||||||
|
// SYN (0x02) + ACK (0x10) = 0x12
|
||||||
|
assert.Equal(t, uint8(0x12), ctx.TCPFlags)
|
||||||
|
assert.Equal(t, uint16(len(data)), ctx.Length)
|
||||||
|
assert.Equal(t, uint8(0), ctx.ICMPType)
|
||||||
|
assert.Equal(t, uint8(0), ctx.ICMPCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPacketContext_IPv4_ICMPTypeCode(t *testing.T) {
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4, TTL: 64, Protocol: layers.IPProtocolICMPv4,
|
||||||
|
SrcIP: net.IPv4(10, 0, 0, 1), DstIP: net.IPv4(10, 0, 0, 2),
|
||||||
|
}
|
||||||
|
// Destination Unreachable, code 3 (port unreachable)
|
||||||
|
icmp := &layers.ICMPv4{
|
||||||
|
TypeCode: layers.CreateICMPv4TypeCode(layers.ICMPv4TypeDestinationUnreachable, layers.ICMPv4CodePort),
|
||||||
|
}
|
||||||
|
data := mustSerialize(t, ip, icmp, gopacket.Payload([]byte{0, 0, 0, 0}))
|
||||||
|
|
||||||
|
var fp firewall.Packet
|
||||||
|
var ctx firewall.PacketContext
|
||||||
|
require.NoError(t, newPacket(data, true, &fp, &ctx))
|
||||||
|
|
||||||
|
assert.Equal(t, uint8(firewall.ProtoICMP), fp.Protocol)
|
||||||
|
assert.Equal(t, uint8(layers.ICMPv4TypeDestinationUnreachable), ctx.ICMPType)
|
||||||
|
assert.Equal(t, uint8(layers.ICMPv4CodePort), ctx.ICMPCode)
|
||||||
|
assert.Equal(t, uint16(len(data)), ctx.Length)
|
||||||
|
assert.Equal(t, uint8(0), ctx.TCPFlags)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPacketContext_IPv4_UDPLengthOnly(t *testing.T) {
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4, TTL: 64, Protocol: layers.IPProtocolUDP,
|
||||||
|
SrcIP: net.IPv4(10, 0, 0, 1), DstIP: net.IPv4(10, 0, 0, 2),
|
||||||
|
}
|
||||||
|
udp := &layers.UDP{SrcPort: 1234, DstPort: 53}
|
||||||
|
require.NoError(t, udp.SetNetworkLayerForChecksum(ip))
|
||||||
|
data := mustSerialize(t, ip, udp, gopacket.Payload([]byte("query")))
|
||||||
|
|
||||||
|
var fp firewall.Packet
|
||||||
|
var ctx firewall.PacketContext
|
||||||
|
require.NoError(t, newPacket(data, true, &fp, &ctx))
|
||||||
|
|
||||||
|
assert.Equal(t, uint8(firewall.ProtoUDP), fp.Protocol)
|
||||||
|
assert.Equal(t, uint16(len(data)), ctx.Length)
|
||||||
|
assert.Zero(t, ctx.TCPFlags)
|
||||||
|
assert.Zero(t, ctx.ICMPType)
|
||||||
|
assert.Zero(t, ctx.ICMPCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPacketContext_IPv6_TCPFlags(t *testing.T) {
|
||||||
|
ip := &layers.IPv6{
|
||||||
|
Version: 6, HopLimit: 64, NextHeader: layers.IPProtocolTCP,
|
||||||
|
SrcIP: net.ParseIP("fd00::1"), DstIP: net.ParseIP("fd00::2"),
|
||||||
|
}
|
||||||
|
tcp := &layers.TCP{SrcPort: 1234, DstPort: 443, FIN: true, ACK: true}
|
||||||
|
require.NoError(t, tcp.SetNetworkLayerForChecksum(ip))
|
||||||
|
data := mustSerialize(t, ip, tcp, gopacket.Payload([]byte("bye")))
|
||||||
|
|
||||||
|
var fp firewall.Packet
|
||||||
|
var ctx firewall.PacketContext
|
||||||
|
require.NoError(t, newPacket(data, true, &fp, &ctx))
|
||||||
|
|
||||||
|
assert.Equal(t, uint8(firewall.ProtoTCP), fp.Protocol)
|
||||||
|
// FIN (0x01) + ACK (0x10) = 0x11
|
||||||
|
assert.Equal(t, uint8(0x11), ctx.TCPFlags)
|
||||||
|
assert.Equal(t, uint16(len(data)), ctx.Length)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPacketContext_IPv6_ICMPv6TypeCode(t *testing.T) {
|
||||||
|
ip := &layers.IPv6{
|
||||||
|
Version: 6, HopLimit: 64, NextHeader: layers.IPProtocolICMPv6,
|
||||||
|
SrcIP: net.ParseIP("fd00::1"), DstIP: net.ParseIP("fd00::2"),
|
||||||
|
}
|
||||||
|
icmp := &layers.ICMPv6{
|
||||||
|
TypeCode: layers.CreateICMPv6TypeCode(layers.ICMPv6TypeDestinationUnreachable, layers.ICMPv6CodePortUnreachable),
|
||||||
|
}
|
||||||
|
require.NoError(t, icmp.SetNetworkLayerForChecksum(ip))
|
||||||
|
data := mustSerialize(t, ip, icmp, gopacket.Payload([]byte{0, 0, 0, 0, 0, 0, 0, 0}))
|
||||||
|
|
||||||
|
var fp firewall.Packet
|
||||||
|
var ctx firewall.PacketContext
|
||||||
|
require.NoError(t, newPacket(data, true, &fp, &ctx))
|
||||||
|
|
||||||
|
assert.Equal(t, uint8(firewall.ProtoICMPv6), fp.Protocol)
|
||||||
|
assert.Equal(t, uint8(layers.ICMPv6TypeDestinationUnreachable), ctx.ICMPType)
|
||||||
|
assert.Equal(t, uint8(layers.ICMPv6CodePortUnreachable), ctx.ICMPCode)
|
||||||
|
assert.Equal(t, uint16(len(data)), ctx.Length)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestPacketContext_NilOK confirms a nil context pointer is accepted by
|
||||||
|
// newPacket (the hot path may elect not to pass one).
|
||||||
|
func TestPacketContext_NilOK(t *testing.T) {
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4, TTL: 64, Protocol: layers.IPProtocolUDP,
|
||||||
|
SrcIP: net.IPv4(10, 0, 0, 1), DstIP: net.IPv4(10, 0, 0, 2),
|
||||||
|
}
|
||||||
|
udp := &layers.UDP{SrcPort: 1, DstPort: 2}
|
||||||
|
require.NoError(t, udp.SetNetworkLayerForChecksum(ip))
|
||||||
|
data := mustSerialize(t, ip, udp)
|
||||||
|
|
||||||
|
var fp firewall.Packet
|
||||||
|
require.NoError(t, newPacket(data, true, &fp, nil))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestPacketContext_FlowCreateCarriesContext exercises the full Drop -> addConn
|
||||||
|
// -> ReportFlowCreate path with a realistic TCP packet and confirms the
|
||||||
|
// context makes it into the reporter.
|
||||||
|
func TestPacketContext_FlowCreateCarriesContext(t *testing.T) {
|
||||||
|
f := newEventFixture(t)
|
||||||
|
r := &recordingReporter{}
|
||||||
|
f.ctl.SetFirewallEventReporter(r)
|
||||||
|
|
||||||
|
// Hand-construct a matching TCP packet.
|
||||||
|
ctx := firewall.PacketContext{Length: 1500, TCPFlags: 0x12}
|
||||||
|
p := f.p
|
||||||
|
p.Protocol = firewall.ProtoTCP
|
||||||
|
require.NoError(t, f.firewall().Drop(p, ctx, true, f.h, f.cp, nil))
|
||||||
|
|
||||||
|
require.Len(t, r.creates, 1)
|
||||||
|
assert.Equal(t, uint16(1500), r.creates[0].ctx.Length)
|
||||||
|
assert.Equal(t, uint8(0x12), r.creates[0].ctx.TCPFlags)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- benchmarks ------------------------------------------------------------
|
||||||
|
|
||||||
|
// noopReporter is the cheapest possible reporter. Methods discard the event.
|
||||||
|
type noopReporter struct{}
|
||||||
|
|
||||||
|
func (noopReporter) ReportDrop(events.DropEvent) {}
|
||||||
|
func (noopReporter) ReportFlowCreate(events.FlowCreateEvent) {}
|
||||||
|
func (noopReporter) ReportFlowEvict(events.FlowEvictEvent) {}
|
||||||
|
func (noopReporter) ReportRulesReload(events.RulesReloadEvent) {
|
||||||
|
}
|
||||||
|
|
||||||
|
// bufferedReporter demonstrates a realistic zero-alloc reporter: each event
|
||||||
|
// is forwarded to a value-typed channel. The channel send is a memcpy into
|
||||||
|
// the channel's pre-allocated ring buffer -- no heap traffic. A background
|
||||||
|
// goroutine would drain these; the bench skips draining to keep the report
|
||||||
|
// path pure.
|
||||||
|
type bufferedReporter struct {
|
||||||
|
drops chan events.DropEvent
|
||||||
|
flows chan events.FlowCreateEvent
|
||||||
|
evicts chan events.FlowEvictEvent
|
||||||
|
}
|
||||||
|
|
||||||
|
func newBufferedReporter(cap int) *bufferedReporter {
|
||||||
|
return &bufferedReporter{
|
||||||
|
drops: make(chan events.DropEvent, cap),
|
||||||
|
flows: make(chan events.FlowCreateEvent, cap),
|
||||||
|
evicts: make(chan events.FlowEvictEvent, cap),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *bufferedReporter) ReportDrop(e events.DropEvent) {
|
||||||
|
select {
|
||||||
|
case r.drops <- e:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *bufferedReporter) ReportFlowCreate(e events.FlowCreateEvent) {
|
||||||
|
select {
|
||||||
|
case r.flows <- e:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *bufferedReporter) ReportFlowEvict(e events.FlowEvictEvent) {
|
||||||
|
select {
|
||||||
|
case r.evicts <- e:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *bufferedReporter) ReportRulesReload(events.RulesReloadEvent) {}
|
||||||
|
|
||||||
|
// pointerReporter is the anti-pattern: it takes the address of the incoming
|
||||||
|
// event struct, which forces the callee-side copy onto the heap. Kept for
|
||||||
|
// comparison so we can see the alloc cost an unwary reporter would incur.
|
||||||
|
type pointerReporter struct {
|
||||||
|
last *events.DropEvent
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *pointerReporter) ReportDrop(e events.DropEvent) {
|
||||||
|
r.last = &e
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *pointerReporter) ReportFlowCreate(events.FlowCreateEvent) {}
|
||||||
|
func (r *pointerReporter) ReportFlowEvict(events.FlowEvictEvent) {}
|
||||||
|
func (r *pointerReporter) ReportRulesReload(events.RulesReloadEvent) {
|
||||||
|
}
|
||||||
|
|
||||||
|
func newBenchFixture(b *testing.B) *eventFixture {
|
||||||
|
b.Helper()
|
||||||
|
l := test.NewLogger()
|
||||||
|
|
||||||
|
vpnNetworks := new(bart.Lite)
|
||||||
|
vpnNetworks.Insert(netip.MustParsePrefix("1.2.3.0/24"))
|
||||||
|
|
||||||
|
c := &dummyCert{
|
||||||
|
name: "host1",
|
||||||
|
networks: []netip.Prefix{netip.MustParsePrefix("1.2.3.4/24")},
|
||||||
|
groups: []string{"default-group"},
|
||||||
|
issuer: "signer-shasum",
|
||||||
|
}
|
||||||
|
h := &HostInfo{
|
||||||
|
ConnectionState: &ConnectionState{
|
||||||
|
peerCert: &cert.CachedCertificate{
|
||||||
|
Certificate: c,
|
||||||
|
InvertedGroups: map[string]struct{}{"default-group": {}},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
vpnAddrs: []netip.Addr{netip.MustParseAddr("1.2.3.4")},
|
||||||
|
}
|
||||||
|
h.buildNetworks(vpnNetworks, c)
|
||||||
|
|
||||||
|
fw := NewFirewall(l, time.Minute, time.Minute, time.Minute, c)
|
||||||
|
// Inbound rule that matches our packet; outbound has no match so we can
|
||||||
|
// also benchmark the no-rule drop path.
|
||||||
|
if err := fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", ""); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctl := &Control{f: &Interface{firewall: fw}, l: l}
|
||||||
|
return &eventFixture{
|
||||||
|
ctl: ctl,
|
||||||
|
fw: fw,
|
||||||
|
p: firewall.Packet{
|
||||||
|
LocalAddr: netip.MustParseAddr("1.2.3.4"),
|
||||||
|
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||||
|
LocalPort: 10,
|
||||||
|
RemotePort: 90,
|
||||||
|
Protocol: firewall.ProtoUDP,
|
||||||
|
},
|
||||||
|
h: h,
|
||||||
|
cp: cert.NewCAPool(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkFirewallDropPath measures the cost of Firewall.Drop on a packet
|
||||||
|
// that reaches the no-matching-rule branch (the longest drop path). Compare
|
||||||
|
// reporter shapes:
|
||||||
|
//
|
||||||
|
// nilReporter -- no reporter installed (feature cost when off)
|
||||||
|
// noopReporter -- reporter installed, methods discard args (minimum on-cost)
|
||||||
|
// bufferedReporter -- realistic zero-alloc reporter: value-typed channels
|
||||||
|
// pointerReporter -- anti-pattern that takes &composite-literal (allocates)
|
||||||
|
func BenchmarkFirewallDropPath(b *testing.B) {
|
||||||
|
run := func(b *testing.B, install func(*Control)) {
|
||||||
|
f := newBenchFixture(b)
|
||||||
|
install(f.ctl)
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
_ = f.firewall().Drop(f.p, firewall.PacketContext{}, false, f.h, f.cp, nil)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
b.Run("nilReporter", func(b *testing.B) { run(b, func(*Control) {}) })
|
||||||
|
b.Run("noopReporter", func(b *testing.B) {
|
||||||
|
run(b, func(c *Control) { c.SetFirewallEventReporter(noopReporter{}) })
|
||||||
|
})
|
||||||
|
b.Run("bufferedReporter", func(b *testing.B) {
|
||||||
|
run(b, func(c *Control) { c.SetFirewallEventReporter(newBufferedReporter(1024)) })
|
||||||
|
})
|
||||||
|
b.Run("pointerReporter", func(b *testing.B) {
|
||||||
|
run(b, func(c *Control) { c.SetFirewallEventReporter(&pointerReporter{}) })
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkConntrackCreate measures Firewall.Drop for an allowed inbound
|
||||||
|
// packet on a fresh conntrack (so addConn fires each iteration).
|
||||||
|
func BenchmarkConntrackCreate(b *testing.B) {
|
||||||
|
run := func(b *testing.B, install func(*Control)) {
|
||||||
|
f := newBenchFixture(b)
|
||||||
|
install(f.ctl)
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
resetConntrack(f.firewall())
|
||||||
|
_ = f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
b.Run("nilReporter", func(b *testing.B) { run(b, func(*Control) {}) })
|
||||||
|
b.Run("noopReporter", func(b *testing.B) {
|
||||||
|
run(b, func(c *Control) { c.SetFirewallEventReporter(noopReporter{}) })
|
||||||
|
})
|
||||||
|
b.Run("bufferedReporter", func(b *testing.B) {
|
||||||
|
run(b, func(c *Control) { c.SetFirewallEventReporter(newBufferedReporter(1024)) })
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkConntrackHit measures the hot path where a flow is already in
|
||||||
|
// conntrack and short-circuits rule evaluation. The reporter slot is checked
|
||||||
|
// only on create/evict, so this bench should show the reporter having zero
|
||||||
|
// impact regardless of install state.
|
||||||
|
func BenchmarkConntrackHit(b *testing.B) {
|
||||||
|
b.Run("nilReporter", func(b *testing.B) {
|
||||||
|
f := newBenchFixture(b)
|
||||||
|
// Prime conntrack.
|
||||||
|
require.NoError(b, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
_ = f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
b.Run("noopReporter", func(b *testing.B) {
|
||||||
|
f := newBenchFixture(b)
|
||||||
|
f.ctl.SetFirewallEventReporter(noopReporter{})
|
||||||
|
require.NoError(b, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil))
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
_ = f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
+108
-320
@@ -3,13 +3,13 @@ package nebula
|
|||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"errors"
|
"errors"
|
||||||
"log/slog"
|
|
||||||
"math"
|
"math"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/cert"
|
"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"
|
||||||
@@ -58,8 +58,9 @@ func TestNewFirewall(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_AddRule(t *testing.T) {
|
func TestFirewall_AddRule(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
|
|
||||||
c := &dummyCert{}
|
c := &dummyCert{}
|
||||||
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||||
@@ -176,8 +177,9 @@ func TestFirewall_AddRule(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_Drop(t *testing.T) {
|
func TestFirewall_Drop(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
myVpnNetworksTable := new(bart.Lite)
|
myVpnNetworksTable := new(bart.Lite)
|
||||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
||||||
p := firewall.Packet{
|
p := firewall.Packet{
|
||||||
@@ -211,49 +213,50 @@ func TestFirewall_Drop(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, false, &h, cp, nil))
|
assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil))
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
// Allow outbound because conntrack
|
// Allow outbound because conntrack
|
||||||
require.NoError(t, fw.Drop(p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil))
|
||||||
|
|
||||||
// test remote mismatch
|
// test remote mismatch
|
||||||
oldRemote := p.RemoteAddr
|
oldRemote := p.RemoteAddr
|
||||||
p.RemoteAddr = netip.MustParseAddr("1.2.3.10")
|
p.RemoteAddr = netip.MustParseAddr("1.2.3.10")
|
||||||
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrInvalidRemoteIP)
|
assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil), ErrInvalidRemoteIP)
|
||||||
p.RemoteAddr = oldRemote
|
p.RemoteAddr = oldRemote
|
||||||
|
|
||||||
// ensure signer doesn't get in the way of group checks
|
// ensure signer doesn't get in the way of group checks
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
|
||||||
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
|
|
||||||
// test caSha doesn't drop on match
|
// test caSha doesn't drop on match
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
|
||||||
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
|
|
||||||
// ensure ca name doesn't get in the way of group checks
|
// ensure ca name doesn't get in the way of group checks
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
|
||||||
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
|
|
||||||
// test caName doesn't drop on match
|
// test caName doesn't drop on match
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
|
||||||
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_DropV6(t *testing.T) {
|
func TestFirewall_DropV6(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
|
|
||||||
myVpnNetworksTable := new(bart.Lite)
|
myVpnNetworksTable := new(bart.Lite)
|
||||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("fd00::/7"))
|
myVpnNetworksTable.Insert(netip.MustParsePrefix("fd00::/7"))
|
||||||
@@ -289,44 +292,44 @@ func TestFirewall_DropV6(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, false, &h, cp, nil))
|
assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil))
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
// Allow outbound because conntrack
|
// Allow outbound because conntrack
|
||||||
require.NoError(t, fw.Drop(p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil))
|
||||||
|
|
||||||
// test remote mismatch
|
// test remote mismatch
|
||||||
oldRemote := p.RemoteAddr
|
oldRemote := p.RemoteAddr
|
||||||
p.RemoteAddr = netip.MustParseAddr("fd12::56")
|
p.RemoteAddr = netip.MustParseAddr("fd12::56")
|
||||||
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrInvalidRemoteIP)
|
assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil), ErrInvalidRemoteIP)
|
||||||
p.RemoteAddr = oldRemote
|
p.RemoteAddr = oldRemote
|
||||||
|
|
||||||
// ensure signer doesn't get in the way of group checks
|
// ensure signer doesn't get in the way of group checks
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
|
||||||
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
|
|
||||||
// test caSha doesn't drop on match
|
// test caSha doesn't drop on match
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
|
||||||
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
|
|
||||||
// ensure ca name doesn't get in the way of group checks
|
// ensure ca name doesn't get in the way of group checks
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
|
||||||
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
|
|
||||||
// test caName doesn't drop on match
|
// test caName doesn't drop on match
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
|
||||||
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
func BenchmarkFirewallTable_match(b *testing.B) {
|
func BenchmarkFirewallTable_match(b *testing.B) {
|
||||||
@@ -482,8 +485,9 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_Drop2(t *testing.T) {
|
func TestFirewall_Drop2(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
myVpnNetworksTable := new(bart.Lite)
|
myVpnNetworksTable := new(bart.Lite)
|
||||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
||||||
|
|
||||||
@@ -533,15 +537,16 @@ func TestFirewall_Drop2(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// h1/c1 lacks the proper groups
|
// h1/c1 lacks the proper groups
|
||||||
require.ErrorIs(t, fw.Drop(p, true, &h1, cp, nil), ErrNoMatchingRule)
|
require.ErrorIs(t, fw.Drop(p, firewall.PacketContext{}, true, &h1, cp, nil), ErrNoMatchingRule)
|
||||||
// c has the proper groups
|
// c has the proper groups
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_Drop3(t *testing.T) {
|
func TestFirewall_Drop3(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
myVpnNetworksTable := new(bart.Lite)
|
myVpnNetworksTable := new(bart.Lite)
|
||||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
||||||
|
|
||||||
@@ -613,23 +618,24 @@ func TestFirewall_Drop3(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// c1 should pass because host match
|
// c1 should pass because host match
|
||||||
require.NoError(t, fw.Drop(p, true, &h1, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h1, cp, nil))
|
||||||
// c2 should pass because ca sha match
|
// c2 should pass because ca sha match
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p, true, &h2, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h2, cp, nil))
|
||||||
// c3 should fail because no match
|
// c3 should fail because no match
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
assert.Equal(t, fw.Drop(p, true, &h3, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, true, &h3, cp, nil), ErrNoMatchingRule)
|
||||||
|
|
||||||
// Test a remote address match
|
// Test a remote address match
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "1.2.3.4/24", "", "", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "1.2.3.4/24", "", "", ""))
|
||||||
require.NoError(t, fw.Drop(p, true, &h1, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h1, cp, nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_Drop3V6(t *testing.T) {
|
func TestFirewall_Drop3V6(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
myVpnNetworksTable := new(bart.Lite)
|
myVpnNetworksTable := new(bart.Lite)
|
||||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("fd00::/7"))
|
myVpnNetworksTable.Insert(netip.MustParsePrefix("fd00::/7"))
|
||||||
|
|
||||||
@@ -661,12 +667,13 @@ func TestFirewall_Drop3V6(t *testing.T) {
|
|||||||
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "fd12::34/120", "", "", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "fd12::34/120", "", "", ""))
|
||||||
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_DropConntrackReload(t *testing.T) {
|
func TestFirewall_DropConntrackReload(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
myVpnNetworksTable := new(bart.Lite)
|
myVpnNetworksTable := new(bart.Lite)
|
||||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
||||||
|
|
||||||
@@ -702,12 +709,12 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
// Allow outbound because conntrack
|
// Allow outbound because conntrack
|
||||||
require.NoError(t, fw.Drop(p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil))
|
||||||
|
|
||||||
oldFw := fw
|
oldFw := fw
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||||
@@ -716,7 +723,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
|||||||
fw.rulesVersion = oldFw.rulesVersion + 1
|
fw.rulesVersion = oldFw.rulesVersion + 1
|
||||||
|
|
||||||
// Allow outbound because conntrack and new rules allow port 10
|
// Allow outbound because conntrack and new rules allow port 10
|
||||||
require.NoError(t, fw.Drop(p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil))
|
||||||
|
|
||||||
oldFw = fw
|
oldFw = fw
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||||
@@ -725,12 +732,13 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
|||||||
fw.rulesVersion = oldFw.rulesVersion + 1
|
fw.rulesVersion = oldFw.rulesVersion + 1
|
||||||
|
|
||||||
// Drop outbound because conntrack doesn't match new ruleset
|
// Drop outbound because conntrack doesn't match new ruleset
|
||||||
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
myVpnNetworksTable := new(bart.Lite)
|
myVpnNetworksTable := new(bart.Lite)
|
||||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
||||||
|
|
||||||
@@ -770,12 +778,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0
|
p.LocalPort = 0
|
||||||
p.RemotePort = 0
|
p.RemotePort = 0
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(*p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
require.NoError(t, fw.Drop(*p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil))
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("nonzero ports", func(t *testing.T) {
|
t.Run("nonzero ports", func(t *testing.T) {
|
||||||
@@ -783,12 +791,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0xabcd
|
p.LocalPort = 0xabcd
|
||||||
p.RemotePort = 0x1234
|
p.RemotePort = 0x1234
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(*p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
require.NoError(t, fw.Drop(*p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil))
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -800,12 +808,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0
|
p.LocalPort = 0
|
||||||
p.RemotePort = 0
|
p.RemotePort = 0
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("nonzero ports, still blocked", func(t *testing.T) {
|
t.Run("nonzero ports, still blocked", func(t *testing.T) {
|
||||||
@@ -813,12 +821,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0xabcd
|
p.LocalPort = 0xabcd
|
||||||
p.RemotePort = 0x1234
|
p.RemotePort = 0x1234
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("nonzero, matching ports, still blocked", func(t *testing.T) {
|
t.Run("nonzero, matching ports, still blocked", func(t *testing.T) {
|
||||||
@@ -826,12 +834,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 80
|
p.LocalPort = 80
|
||||||
p.RemotePort = 80
|
p.RemotePort = 80
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
t.Run("Any proto, any port", func(t *testing.T) {
|
t.Run("Any proto, any port", func(t *testing.T) {
|
||||||
@@ -843,12 +851,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0
|
p.LocalPort = 0
|
||||||
p.RemotePort = 0
|
p.RemotePort = 0
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(*p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
require.NoError(t, fw.Drop(*p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil))
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("nonzero ports, allowed", func(t *testing.T) {
|
t.Run("nonzero ports, allowed", func(t *testing.T) {
|
||||||
@@ -857,23 +865,24 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0xabcd
|
p.LocalPort = 0xabcd
|
||||||
p.RemotePort = 0x1234
|
p.RemotePort = 0x1234
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(*p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil))
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
require.NoError(t, fw.Drop(*p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil))
|
||||||
//different ID is blocked
|
//different ID is blocked
|
||||||
p.RemotePort++
|
p.RemotePort++
|
||||||
require.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
require.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_DropIPSpoofing(t *testing.T) {
|
func TestFirewall_DropIPSpoofing(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
myVpnNetworksTable := new(bart.Lite)
|
myVpnNetworksTable := new(bart.Lite)
|
||||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("192.0.2.1/24"))
|
myVpnNetworksTable.Insert(netip.MustParsePrefix("192.0.2.1/24"))
|
||||||
|
|
||||||
@@ -913,160 +922,7 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
|
|||||||
Protocol: firewall.ProtoUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
Fragment: false,
|
Fragment: false,
|
||||||
}
|
}
|
||||||
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
|
assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, 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) {
|
||||||
@@ -1182,101 +1038,32 @@ 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
|
||||||
c := &dummyCert{}
|
c := &dummyCert{}
|
||||||
cs, err := newCertState(cert.Version2, nil, c, false, cert.Curve_CURVE25519, nil, "aes")
|
cs, err := newCertState(cert.Version2, nil, c, false, cert.Curve_CURVE25519, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(l)
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": "asdf"}
|
conf.Settings["firewall"] = map[string]any{"outbound": "asdf"}
|
||||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||||
require.EqualError(t, err, "firewall.outbound failed to parse, should be an array of rules")
|
require.EqualError(t, err, "firewall.outbound failed to parse, should be an array of rules")
|
||||||
|
|
||||||
// Test both port and code
|
// Test both port and code
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "code": "2"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "code": "2"}}}
|
||||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||||
require.EqualError(t, err, "firewall.outbound rule #0; only one of port or code should be provided")
|
require.EqualError(t, err, "firewall.outbound rule #0; only one of port or code should be provided")
|
||||||
|
|
||||||
// Test missing host, group, cidr, ca_name and ca_sha
|
// Test missing host, group, cidr, ca_name and ca_sha
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{}}}
|
||||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||||
require.EqualError(t, err, "firewall.outbound rule #0; at least one of host, group, cidr, local_cidr, ca_name, or ca_sha must be provided")
|
require.EqualError(t, err, "firewall.outbound rule #0; at least one of host, group, cidr, local_cidr, ca_name, or ca_sha must be provided")
|
||||||
|
|
||||||
// Test code/port error
|
// Test code/port error
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "a", "host": "testh", "proto": "any"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "a", "host": "testh", "proto": "any"}}}
|
||||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||||
require.EqualError(t, err, "firewall.outbound rule #0; code was not a number; `a`")
|
require.EqualError(t, err, "firewall.outbound rule #0; code was not a number; `a`")
|
||||||
@@ -1286,25 +1073,25 @@ func TestNewFirewallFromConfig(t *testing.T) {
|
|||||||
require.EqualError(t, err, "firewall.outbound rule #0; port was not a number; `a`")
|
require.EqualError(t, err, "firewall.outbound rule #0; port was not a number; `a`")
|
||||||
|
|
||||||
// Test proto error
|
// Test proto error
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "1", "host": "testh"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "1", "host": "testh"}}}
|
||||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||||
require.EqualError(t, err, "firewall.outbound rule #0; proto was not understood; ``")
|
require.EqualError(t, err, "firewall.outbound rule #0; proto was not understood; ``")
|
||||||
|
|
||||||
// Test cidr parse error
|
// Test cidr parse error
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "1", "cidr": "testh", "proto": "any"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "1", "cidr": "testh", "proto": "any"}}}
|
||||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||||
require.EqualError(t, err, "firewall.outbound rule #0; cidr did not parse; netip.ParsePrefix(\"testh\"): no '/'")
|
require.EqualError(t, err, "firewall.outbound rule #0; cidr did not parse; netip.ParsePrefix(\"testh\"): no '/'")
|
||||||
|
|
||||||
// Test local_cidr parse error
|
// Test local_cidr parse error
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "1", "local_cidr": "testh", "proto": "any"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "1", "local_cidr": "testh", "proto": "any"}}}
|
||||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||||
require.EqualError(t, err, "firewall.outbound rule #0; local_cidr did not parse; netip.ParsePrefix(\"testh\"): no '/'")
|
require.EqualError(t, err, "firewall.outbound rule #0; local_cidr did not parse; netip.ParsePrefix(\"testh\"): no '/'")
|
||||||
|
|
||||||
// Test both group and groups
|
// Test both group and groups
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "group": "a", "groups": []string{"b", "c"}}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "group": "a", "groups": []string{"b", "c"}}}}
|
||||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||||
require.EqualError(t, err, "firewall.inbound rule #0; only one of group or groups should be defined, both provided")
|
require.EqualError(t, err, "firewall.inbound rule #0; only one of group or groups should be defined, both provided")
|
||||||
@@ -1313,35 +1100,35 @@ func TestNewFirewallFromConfig(t *testing.T) {
|
|||||||
func TestAddFirewallRulesFromConfig(t *testing.T) {
|
func TestAddFirewallRulesFromConfig(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
// Test adding tcp rule
|
// Test adding tcp rule
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(l)
|
||||||
mf := &mockFirewall{}
|
mf := &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "tcp", "host": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "tcp", "host": "a"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding udp rule
|
// Test adding udp rule
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "udp", "host": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "udp", "host": "a"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding icmp rule
|
// Test adding icmp rule
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "icmp", "host": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "icmp", "host": "a"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding icmp rule no port
|
// Test adding icmp rule no port
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"proto": "icmp", "host": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"proto": "icmp", "host": "a"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding any rule
|
// Test adding any rule
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "host": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "host": "a"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
@@ -1349,14 +1136,14 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
|
|||||||
|
|
||||||
// Test adding rule with cidr
|
// Test adding rule with cidr
|
||||||
cidr := netip.MustParsePrefix("10.0.0.0/8")
|
cidr := netip.MustParsePrefix("10.0.0.0/8")
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": cidr.String()}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": cidr.String()}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: cidr.String(), localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: cidr.String(), localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding rule with local_cidr
|
// Test adding rule with local_cidr
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": cidr.String()}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": cidr.String()}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
@@ -1364,82 +1151,82 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
|
|||||||
|
|
||||||
// Test adding rule with cidr ipv6
|
// Test adding rule with cidr ipv6
|
||||||
cidr6 := netip.MustParsePrefix("fd00::/8")
|
cidr6 := netip.MustParsePrefix("fd00::/8")
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": cidr6.String()}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": cidr6.String()}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: cidr6.String(), localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: cidr6.String(), localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding rule with any cidr
|
// Test adding rule with any cidr
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": "any"}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": "any"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "any", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "any", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding rule with junk cidr
|
// Test adding rule with junk cidr
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": "junk/junk"}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": "junk/junk"}}}
|
||||||
require.EqualError(t, AddFirewallRulesFromConfig(l, true, conf, mf), "firewall.inbound rule #0; cidr did not parse; netip.ParsePrefix(\"junk/junk\"): ParseAddr(\"junk\"): unable to parse IP")
|
require.EqualError(t, AddFirewallRulesFromConfig(l, true, conf, mf), "firewall.inbound rule #0; cidr did not parse; netip.ParsePrefix(\"junk/junk\"): ParseAddr(\"junk\"): unable to parse IP")
|
||||||
|
|
||||||
// Test adding rule with local_cidr ipv6
|
// Test adding rule with local_cidr ipv6
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": cidr6.String()}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": cidr6.String()}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "", localIp: cidr6.String()}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "", localIp: cidr6.String()}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding rule with any local_cidr
|
// Test adding rule with any local_cidr
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": "any"}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": "any"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, localIp: "any"}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, localIp: "any"}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding rule with junk local_cidr
|
// Test adding rule with junk local_cidr
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": "junk/junk"}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": "junk/junk"}}}
|
||||||
require.EqualError(t, AddFirewallRulesFromConfig(l, true, conf, mf), "firewall.inbound rule #0; local_cidr did not parse; netip.ParsePrefix(\"junk/junk\"): ParseAddr(\"junk\"): unable to parse IP")
|
require.EqualError(t, AddFirewallRulesFromConfig(l, true, conf, mf), "firewall.inbound rule #0; local_cidr did not parse; netip.ParsePrefix(\"junk/junk\"): ParseAddr(\"junk\"): unable to parse IP")
|
||||||
|
|
||||||
// Test adding rule with ca_sha
|
// Test adding rule with ca_sha
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "ca_sha": "12312313123"}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "ca_sha": "12312313123"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "", localIp: "", caSha: "12312313123"}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "", localIp: "", caSha: "12312313123"}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding rule with ca_name
|
// Test adding rule with ca_name
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "ca_name": "root01"}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "ca_name": "root01"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "", localIp: "", caName: "root01"}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "", localIp: "", caName: "root01"}, mf.lastCall)
|
||||||
|
|
||||||
// Test single group
|
// Test single group
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "group": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "group": "a"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: []string{"a"}, ip: "", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: []string{"a"}, ip: "", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test single groups
|
// Test single groups
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "groups": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "groups": "a"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: []string{"a"}, ip: "", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: []string{"a"}, ip: "", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test multiple AND groups
|
// Test multiple AND groups
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "groups": []string{"a", "b"}}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "groups": []string{"a", "b"}}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: []string{"a", "b"}, ip: "", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: []string{"a", "b"}, ip: "", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test Add error
|
// Test Add error
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(l)
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
mf.nextCallReturn = errors.New("test error")
|
mf.nextCallReturn = errors.New("test error")
|
||||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "host": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "host": "a"}}}
|
||||||
@@ -1447,8 +1234,9 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_convertRule(t *testing.T) {
|
func TestFirewall_convertRule(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
|
|
||||||
// Ensure group array of 1 is converted and a warning is printed
|
// Ensure group array of 1 is converted and a warning is printed
|
||||||
c := map[string]any{
|
c := map[string]any{
|
||||||
@@ -1456,9 +1244,7 @@ func TestFirewall_convertRule(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
r, err := convertRule(l, c, "test", 1)
|
r, err := convertRule(l, c, "test", 1)
|
||||||
assert.Contains(t, ob.String(), "group was an array with a single value, converting to simple value")
|
assert.Contains(t, ob.String(), "test rule #1; group was an array with a single value, converting to simple value")
|
||||||
assert.Contains(t, ob.String(), "table=test")
|
|
||||||
assert.Contains(t, ob.String(), "rule=1")
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, []string{"group1"}, r.Groups)
|
assert.Equal(t, []string{"group1"}, r.Groups)
|
||||||
|
|
||||||
@@ -1484,8 +1270,9 @@ func TestFirewall_convertRule(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_convertRuleSanity(t *testing.T) {
|
func TestFirewall_convertRuleSanity(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
|
|
||||||
noWarningPlease := []map[string]any{
|
noWarningPlease := []map[string]any{
|
||||||
{"group": "group1"},
|
{"group": "group1"},
|
||||||
@@ -1549,7 +1336,7 @@ func (c *testcase) Test(t *testing.T, fw *Firewall) {
|
|||||||
t.Helper()
|
t.Helper()
|
||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
err := fw.Drop(c.p, true, c.h, cp, nil)
|
err := fw.Drop(c.p, firewall.PacketContext{}, true, c.h, cp, nil)
|
||||||
if c.err == nil {
|
if c.err == nil {
|
||||||
require.NoError(t, err, "failed to not drop remote address %s", c.p.RemoteAddr)
|
require.NoError(t, err, "failed to not drop remote address %s", c.p.RemoteAddr)
|
||||||
} else {
|
} else {
|
||||||
@@ -1599,7 +1386,7 @@ type testsetup struct {
|
|||||||
fw *Firewall
|
fw *Firewall
|
||||||
}
|
}
|
||||||
|
|
||||||
func newSetup(t *testing.T, l *slog.Logger, myPrefixes ...netip.Prefix) testsetup {
|
func newSetup(t *testing.T, l *logrus.Logger, myPrefixes ...netip.Prefix) testsetup {
|
||||||
c := dummyCert{
|
c := dummyCert{
|
||||||
name: "me",
|
name: "me",
|
||||||
networks: myPrefixes,
|
networks: myPrefixes,
|
||||||
@@ -1610,7 +1397,7 @@ func newSetup(t *testing.T, l *slog.Logger, myPrefixes ...netip.Prefix) testsetu
|
|||||||
return newSetupFromCert(t, l, c)
|
return newSetupFromCert(t, l, c)
|
||||||
}
|
}
|
||||||
|
|
||||||
func newSetupFromCert(t *testing.T, l *slog.Logger, c dummyCert) testsetup {
|
func newSetupFromCert(t *testing.T, l *logrus.Logger, c dummyCert) testsetup {
|
||||||
myVpnNetworksTable := new(bart.Lite)
|
myVpnNetworksTable := new(bart.Lite)
|
||||||
for _, prefix := range c.Networks() {
|
for _, prefix := range c.Networks() {
|
||||||
myVpnNetworksTable.Insert(prefix)
|
myVpnNetworksTable.Insert(prefix)
|
||||||
@@ -1627,8 +1414,9 @@ func newSetupFromCert(t *testing.T, l *slog.Logger, c dummyCert) testsetup {
|
|||||||
|
|
||||||
func TestFirewall_Drop_EnforceIPMatch(t *testing.T) {
|
func TestFirewall_Drop_EnforceIPMatch(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
l := test.NewLogger()
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutput(ob)
|
l.SetOutput(ob)
|
||||||
|
|
||||||
myPrefix := netip.MustParsePrefix("1.1.1.1/8")
|
myPrefix := netip.MustParsePrefix("1.1.1.1/8")
|
||||||
// for now, it's okay that these are all "incoming", the logic this test tries to check doesn't care about in/out
|
// for now, it's okay that these are all "incoming", the logic this test tries to check doesn't care about in/out
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
module github.com/slackhq/nebula
|
module github.com/slackhq/nebula
|
||||||
|
|
||||||
go 1.26.0
|
go 1.25.0
|
||||||
|
|
||||||
require (
|
require (
|
||||||
dario.cat/mergo v1.0.2
|
dario.cat/mergo v1.0.2
|
||||||
@@ -9,30 +9,30 @@ 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.28.0
|
github.com/gaissmai/bart v0.26.0
|
||||||
github.com/gogo/protobuf v1.3.2
|
github.com/gogo/protobuf v1.3.2
|
||||||
github.com/google/gopacket v1.1.19
|
github.com/google/gopacket v1.1.19
|
||||||
github.com/kardianos/service v1.3.0
|
github.com/kardianos/service v1.2.4
|
||||||
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
|
||||||
github.com/prometheus/client_golang v1.23.2
|
github.com/prometheus/client_golang v1.23.2
|
||||||
github.com/rcrowley/go-metrics v0.0.0-20201227073835-cf1acfcdf475
|
github.com/rcrowley/go-metrics v0.0.0-20201227073835-cf1acfcdf475
|
||||||
|
github.com/sirupsen/logrus v1.9.4
|
||||||
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e
|
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e
|
||||||
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
||||||
github.com/stretchr/testify v1.11.1
|
github.com/stretchr/testify v1.11.1
|
||||||
github.com/vishvananda/netlink v1.3.1
|
github.com/vishvananda/netlink v1.3.1
|
||||||
go.uber.org/goleak v1.3.0
|
|
||||||
go.yaml.in/yaml/v3 v3.0.4
|
go.yaml.in/yaml/v3 v3.0.4
|
||||||
golang.org/x/crypto v0.54.0
|
golang.org/x/crypto v0.50.0
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
||||||
golang.org/x/net v0.57.0
|
golang.org/x/net v0.52.0
|
||||||
golang.org/x/sync v0.22.0
|
golang.org/x/sync v0.20.0
|
||||||
golang.org/x/sys v0.47.0
|
golang.org/x/sys v0.43.0
|
||||||
golang.org/x/term v0.45.0
|
golang.org/x/term v0.42.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 v1.0.1
|
golang.zx2c4.com/wireguard/windows v0.6.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.36.0 // indirect
|
golang.org/x/mod v0.34.0 // indirect
|
||||||
golang.org/x/time v0.5.0 // indirect
|
golang.org/x/time v0.5.0 // indirect
|
||||||
golang.org/x/tools v0.45.0 // indirect
|
golang.org/x/tools v0.43.0 // indirect
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -26,8 +26,8 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c
|
|||||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/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.28.0 h1:89yZLo8NmyqD0RYgJ3QO9HhqqGGw+oWhf90cZm69Lko=
|
github.com/gaissmai/bart v0.26.0 h1:xOZ57E9hJLBiQaSyeZa9wgWhGuzfGACgqp4BE77OkO0=
|
||||||
github.com/gaissmai/bart v0.28.0/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c=
|
github.com/gaissmai/bart v0.26.0/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c=
|
||||||
github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||||
github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||||
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
||||||
@@ -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.3.0 h1:/LGy+xPP2TM+GLTiCZ2di7cy0Jd/qrawlTUfqKYFdTI=
|
github.com/kardianos/service v1.2.4 h1:XNlGtZOYNx2u91urOdg/Kfmc+gfmuIo1Dd3rEi2OgBk=
|
||||||
github.com/kardianos/service v1.3.0/go.mod h1:E4V9ufUuY82F7Ztlu1eN9VXWIQxg8NoLQlmFe0MtrXc=
|
github.com/kardianos/service v1.2.4/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=
|
||||||
@@ -133,6 +133,8 @@ github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncj
|
|||||||
github.com/sirupsen/logrus v1.2.0/go.mod h1:LxeOpSwHxABJmUn/MG1IvRgCAasNZTLOkJPxbbu5VWo=
|
github.com/sirupsen/logrus v1.2.0/go.mod h1:LxeOpSwHxABJmUn/MG1IvRgCAasNZTLOkJPxbbu5VWo=
|
||||||
github.com/sirupsen/logrus v1.4.2/go.mod h1:tLMulIdttU9McNUspp0xgXVQah82FyeX6MwdIuYE2rE=
|
github.com/sirupsen/logrus v1.4.2/go.mod h1:tLMulIdttU9McNUspp0xgXVQah82FyeX6MwdIuYE2rE=
|
||||||
github.com/sirupsen/logrus v1.6.0/go.mod h1:7uNnSEd1DgxDLC74fIahvMZmmYsHGZGEOFrfsX/uA88=
|
github.com/sirupsen/logrus v1.6.0/go.mod h1:7uNnSEd1DgxDLC74fIahvMZmmYsHGZGEOFrfsX/uA88=
|
||||||
|
github.com/sirupsen/logrus v1.9.4 h1:TsZE7l11zFCLZnZ+teH4Umoq5BhEIfIzfRDZ1Uzql2w=
|
||||||
|
github.com/sirupsen/logrus v1.9.4/go.mod h1:ftWc9WdOfJ0a92nsE2jF5u5ZwH8Bv2zdeOC42RjbV2g=
|
||||||
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e h1:MRM5ITcdelLK2j1vwZ3Je0FKVCfqOLp5zO6trqMLYs0=
|
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e h1:MRM5ITcdelLK2j1vwZ3Je0FKVCfqOLp5zO6trqMLYs0=
|
||||||
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e/go.mod h1:XV66xRDqSt+GTGFMVlhk3ULuV0y9ZmzeVGR4mloJI3M=
|
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e/go.mod h1:XV66xRDqSt+GTGFMVlhk3ULuV0y9ZmzeVGR4mloJI3M=
|
||||||
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6 h1:pnnLyeX7o/5aX8qUQ69P/mLojDqwda8hFOCBTmP/6hw=
|
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6 h1:pnnLyeX7o/5aX8qUQ69P/mLojDqwda8hFOCBTmP/6hw=
|
||||||
@@ -162,16 +164,16 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
|
|||||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
golang.org/x/crypto v0.0.0-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.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
|
golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
|
||||||
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
|
golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 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.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4=
|
golang.org/x/mod v0.34.0 h1:xIHgNUUnW6sYkcM5Jleh05DvLOtwc6RitGHbDk4akRI=
|
||||||
golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ=
|
golang.org/x/mod v0.34.0/go.mod h1:ykgH52iCZe79kzLLMhyCUzhMci+nQj+0XkbXpNYtVjY=
|
||||||
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
golang.org/x/net v0.0.0-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 +184,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
|
|||||||
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
golang.org/x/net v0.0.0-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.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
|
golang.org/x/net v0.52.0 h1:He/TN1l0e4mmR3QqHMT2Xab3Aj3L9qjbhRm78/6jrW0=
|
||||||
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
|
golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw=
|
||||||
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
||||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
@@ -191,8 +193,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
|||||||
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-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.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
||||||
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||||
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-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 +210,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w
|
|||||||
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.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.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI=
|
||||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||||
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
|
golang.org/x/term v0.42.0 h1:UiKe+zDFmJobeJ5ggPwOshJIVt6/Ft0rcfrXZDLWAWY=
|
||||||
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
|
golang.org/x/term v0.42.0/go.mod h1:Dq/D+snpsbazcBG5+F9Q1n2rXV8Ma+71xEjTRufARgY=
|
||||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
golang.org/x/text v0.3.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 +225,8 @@ golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtn
|
|||||||
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
golang.org/x/tools v0.0.0-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.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8=
|
golang.org/x/tools v0.43.0 h1:12BdW9CeB3Z+J/I/wj34VMl8X+fEXBxVR90JeMX5E7s=
|
||||||
golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0=
|
golang.org/x/tools v0.43.0/go.mod h1:uHkMso649BX2cZK6+RpuIPXS3ho2hZo4FVwfoy1vIk0=
|
||||||
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-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 +235,8 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu
|
|||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
golang.zx2c4.com/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 v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8=
|
golang.zx2c4.com/wireguard/windows v0.6.1 h1:XMaKojH1Hs/raMrmnir4n35nTvzvWj7NmSYzHn2F4qU=
|
||||||
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
|
golang.zx2c4.com/wireguard/windows v0.6.1/go.mod h1:04aqInu5GYuTFvMuDw/rKBAF7mHrltW/3rekpfbbZDM=
|
||||||
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
|
google.golang.org/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=
|
||||||
|
|||||||
@@ -1,57 +0,0 @@
|
|||||||
package handshake
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/rand"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Credential holds everything needed to participate in a handshake
|
|
||||||
// at a given cert version. Version and Curve are read from Cert; the public
|
|
||||||
// half of the static keypair likewise comes from Cert.PublicKey().
|
|
||||||
type Credential struct {
|
|
||||||
Cert cert.Certificate // the certificate
|
|
||||||
Bytes []byte // pre-marshaled certificate bytes
|
|
||||||
privateKey []byte // static private key (public half lives in Cert)
|
|
||||||
cipherSuite noise.CipherSuite // pre-built cipher suite (DH + cipher + hash)
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewCredential creates a Credential with all material needed for handshake
|
|
||||||
// participation. The cipherSuite should be pre-built by the caller with the
|
|
||||||
// appropriate DH function, cipher, and hash.
|
|
||||||
func NewCredential(
|
|
||||||
c cert.Certificate,
|
|
||||||
hsBytes []byte,
|
|
||||||
privateKey []byte,
|
|
||||||
cipherSuite noise.CipherSuite,
|
|
||||||
) *Credential {
|
|
||||||
return &Credential{
|
|
||||||
Cert: c,
|
|
||||||
Bytes: hsBytes,
|
|
||||||
privateKey: privateKey,
|
|
||||||
cipherSuite: cipherSuite,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildHandshakeState creates a noise.HandshakeState from this credential.
|
|
||||||
func (hc *Credential) buildHandshakeState(initiator bool, pattern noise.HandshakePattern) (*noise.HandshakeState, error) {
|
|
||||||
return noise.NewHandshakeState(noise.Config{
|
|
||||||
CipherSuite: hc.cipherSuite,
|
|
||||||
Random: rand.Reader,
|
|
||||||
Pattern: pattern,
|
|
||||||
Initiator: initiator,
|
|
||||||
StaticKeypair: noise.DHKey{Private: hc.privateKey, Public: hc.Cert.PublicKey()},
|
|
||||||
PresharedKey: []byte{},
|
|
||||||
PresharedKeyPlacement: 0,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetCredentialFunc returns the handshake credential for the given version,
|
|
||||||
// or nil if that version is not available.
|
|
||||||
//
|
|
||||||
// Implementations must return credentials drawn from a snapshot stable for
|
|
||||||
// the lifetime of any single Machine. The Machine may call this multiple
|
|
||||||
// times during a handshake (e.g. when negotiating to the peer's version)
|
|
||||||
// and assumes the underlying static keypair is consistent across calls.
|
|
||||||
type GetCredentialFunc func(v cert.Version) *Credential
|
|
||||||
@@ -1,22 +0,0 @@
|
|||||||
package handshake
|
|
||||||
|
|
||||||
import "errors"
|
|
||||||
|
|
||||||
var (
|
|
||||||
ErrInitiateOnResponder = errors.New("initiate called on responder")
|
|
||||||
ErrInitiateAlreadyCalled = errors.New("initiate already called")
|
|
||||||
ErrInitiateNotCalled = errors.New("initiate must be called before ProcessPacket for initiators")
|
|
||||||
ErrPacketTooShort = errors.New("packet too short")
|
|
||||||
ErrPublicKeyMismatch = errors.New("public key mismatch between certificate and handshake")
|
|
||||||
ErrIncompleteHandshake = errors.New("handshake completed without receiving required content")
|
|
||||||
ErrMachineFailed = errors.New("handshake machine has failed")
|
|
||||||
ErrUnknownSubtype = errors.New("unknown handshake subtype")
|
|
||||||
ErrMissingContent = errors.New("expected handshake content but message was empty")
|
|
||||||
ErrUnexpectedContent = errors.New("received unexpected handshake content")
|
|
||||||
ErrInvalidRemoteIndex = errors.New("peer sent an invalid index in handshake payload")
|
|
||||||
ErrIndexAllocation = errors.New("failed to allocate local index")
|
|
||||||
ErrNoCredential = errors.New("no handshake credential available for cert version")
|
|
||||||
ErrAsymmetricCipherKeys = errors.New("noise produced only one cipher key")
|
|
||||||
ErrMultiMessageUnsupported = errors.New("multi-message handshake patterns are not yet supported by the manager")
|
|
||||||
ErrSubtypeMismatch = errors.New("packet subtype does not match handshake machine subtype")
|
|
||||||
)
|
|
||||||
@@ -1,29 +0,0 @@
|
|||||||
// This file documents the wire format the nebula handshake speaks. It is
|
|
||||||
// not run through protoc; the encoder/decoder in payload.go is hand-written
|
|
||||||
// against this shape directly to keep the parser narrow and panic-free.
|
|
||||||
//
|
|
||||||
// Any change to the wire format must be reflected here, and adding a new
|
|
||||||
// field requires updating MarshalPayload / unmarshalPayloadDetails together
|
|
||||||
// with the field-uniqueness and wire-type checks in those functions.
|
|
||||||
|
|
||||||
syntax = "proto3";
|
|
||||||
package nebula.handshake;
|
|
||||||
|
|
||||||
message NebulaHandshake {
|
|
||||||
NebulaHandshakeDetails Details = 1;
|
|
||||||
bytes Hmac = 2;
|
|
||||||
}
|
|
||||||
|
|
||||||
message NebulaHandshakeDetails {
|
|
||||||
bytes Cert = 1;
|
|
||||||
uint32 InitiatorIndex = 2;
|
|
||||||
uint32 ResponderIndex = 3;
|
|
||||||
// Cookie was reserved for an anti-DoS mechanism that was never
|
|
||||||
// implemented. No released version of nebula has ever populated it; the
|
|
||||||
// hand-written parser silently skips it on read.
|
|
||||||
uint64 Cookie = 4 [deprecated = true];
|
|
||||||
uint64 Time = 5;
|
|
||||||
uint32 CertVersion = 8;
|
|
||||||
// reserved for WIP multiport
|
|
||||||
reserved 6, 7;
|
|
||||||
}
|
|
||||||
@@ -1,116 +0,0 @@
|
|||||||
package handshake
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
ct "github.com/slackhq/nebula/cert_test"
|
|
||||||
"github.com/slackhq/nebula/header"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
// testCertState holds cert material for a test peer.
|
|
||||||
type testCertState struct {
|
|
||||||
version cert.Version
|
|
||||||
creds map[cert.Version]*Credential
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *testCertState) getCredential(v cert.Version) *Credential {
|
|
||||||
return s.creds[v]
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTestCertState(
|
|
||||||
t *testing.T, ca cert.Certificate, caKey []byte, name string, networks []netip.Prefix,
|
|
||||||
) *testCertState {
|
|
||||||
return newTestCertStateWithCipher(t, ca, caKey, name, networks, noise.CipherChaChaPoly)
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTestCertStateWithCipher(
|
|
||||||
t *testing.T, ca cert.Certificate, caKey []byte, name string, networks []netip.Prefix,
|
|
||||||
cipher noise.CipherFunc,
|
|
||||||
) *testCertState {
|
|
||||||
t.Helper()
|
|
||||||
c, _, rawPrivKey, _ := ct.NewTestCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, ca, caKey,
|
|
||||||
name, ca.NotBefore(), ca.NotAfter(), networks, nil, nil,
|
|
||||||
)
|
|
||||||
|
|
||||||
priv, _, _, err := cert.UnmarshalPrivateKeyFromPEM(rawPrivKey)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
hsBytes, err := c.MarshalForHandshakes()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
ncs := noise.NewCipherSuite(noise.DH25519, cipher, noise.HashSHA256)
|
|
||||||
return &testCertState{
|
|
||||||
version: cert.Version2,
|
|
||||||
creds: map[cert.Version]*Credential{
|
|
||||||
cert.Version2: NewCredential(c, hsBytes, priv, ncs),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func testVerifier(pool *cert.CAPool) CertVerifier {
|
|
||||||
return func(c cert.Certificate) (*cert.CachedCertificate, error) {
|
|
||||||
return pool.VerifyCertificate(time.Now(), c)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTestMachine(
|
|
||||||
t *testing.T,
|
|
||||||
cs *testCertState,
|
|
||||||
verifier CertVerifier,
|
|
||||||
initiator bool,
|
|
||||||
localIndex uint32,
|
|
||||||
) *Machine {
|
|
||||||
t.Helper()
|
|
||||||
m, err := NewMachine(
|
|
||||||
cs.version, cs.getCredential,
|
|
||||||
verifier, func() (uint32, error) { return localIndex, nil },
|
|
||||||
initiator, header.HandshakeIXPSK0,
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
return m
|
|
||||||
}
|
|
||||||
|
|
||||||
func initiateHandshake(
|
|
||||||
t *testing.T,
|
|
||||||
initCS *testCertState, initVerifier CertVerifier,
|
|
||||||
respCS *testCertState, respVerifier CertVerifier,
|
|
||||||
) (initM, respM *Machine, respResult *Result, resp []byte, err error) {
|
|
||||||
t.Helper()
|
|
||||||
initM = newTestMachine(t, initCS, initVerifier, true, 100)
|
|
||||||
msg1, merr := initM.Initiate(nil)
|
|
||||||
require.NoError(t, merr)
|
|
||||||
|
|
||||||
respM = newTestMachine(t, respCS, respVerifier, false, 200)
|
|
||||||
resp, respResult, err = respM.ProcessPacket(nil, msg1)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
func doFullHandshake(
|
|
||||||
t *testing.T, initCS, respCS *testCertState, caPool *cert.CAPool,
|
|
||||||
) (initResult, respResult *Result) {
|
|
||||||
t.Helper()
|
|
||||||
v := testVerifier(caPool)
|
|
||||||
|
|
||||||
initM := newTestMachine(t, initCS, v, true, 1000)
|
|
||||||
respM := newTestMachine(t, respCS, v, false, 2000)
|
|
||||||
|
|
||||||
msg1, err := initM.Initiate(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
resp, respResult, err := respM.ProcessPacket(nil, msg1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, respResult)
|
|
||||||
require.NotEmpty(t, resp)
|
|
||||||
|
|
||||||
_, initResult, err = initM.ProcessPacket(nil, resp)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, initResult)
|
|
||||||
|
|
||||||
return initResult, respResult
|
|
||||||
}
|
|
||||||
@@ -1,454 +0,0 @@
|
|||||||
package handshake
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"fmt"
|
|
||||||
"slices"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/header"
|
|
||||||
)
|
|
||||||
|
|
||||||
// IndexAllocator is called by the Machine to allocate a local index for the
|
|
||||||
// handshake. It is called at most once, when the first outgoing message that
|
|
||||||
// carries a payload is built.
|
|
||||||
//
|
|
||||||
// Implementations MUST NOT return 0. Zero is reserved as a sentinel meaning
|
|
||||||
// "no index assigned" on the wire and in the payload-presence checks. If an
|
|
||||||
// allocator ever returned 0, a legitimate handshake's payload could be
|
|
||||||
// indistinguishable from an empty one and would be rejected.
|
|
||||||
type IndexAllocator func() (uint32, error)
|
|
||||||
|
|
||||||
// CertVerifier is called by the Machine after reconstructing the peer's
|
|
||||||
// certificate from the handshake. The verifier performs all validation
|
|
||||||
// (CA trust, expiry, policy checks, allow lists).
|
|
||||||
type CertVerifier func(cert.Certificate) (*cert.CachedCertificate, error)
|
|
||||||
|
|
||||||
// Result contains the results of a successful handshake.
|
|
||||||
// Returned by ProcessPacket when the handshake is complete.
|
|
||||||
type Result struct {
|
|
||||||
EKey *noise.CipherState
|
|
||||||
DKey *noise.CipherState
|
|
||||||
Cipher noise.CipherFunc // identifies which post-handshake CipherState the data plane should wrap EKey/DKey in
|
|
||||||
MyCert cert.Certificate
|
|
||||||
RemoteCert *cert.CachedCertificate
|
|
||||||
RemoteIndex uint32
|
|
||||||
LocalIndex uint32
|
|
||||||
HandshakeTime uint64
|
|
||||||
MessageIndex uint64 // number of messages exchanged during the handshake
|
|
||||||
Initiator bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// Machine drives a Noise handshake through N messages. It handles Noise
|
|
||||||
// protocol operations, certificate reconstruction, and payload encoding.
|
|
||||||
// Certificate validation is delegated to the caller via CertVerifier.
|
|
||||||
//
|
|
||||||
// A Machine is not safe for concurrent use. The caller must ensure that
|
|
||||||
// Initiate and ProcessPacket are not called concurrently.
|
|
||||||
//
|
|
||||||
// Error contract: when ProcessPacket or Initiate returns an error, callers
|
|
||||||
// must check Failed() to decide what to do next. If Failed() is false the
|
|
||||||
// underlying noise state was not advanced (the packet was rejected before
|
|
||||||
// ReadMessage took effect, or the rejection is non-fatal like a stale
|
|
||||||
// retransmit) and the Machine can accept another packet. If Failed() is
|
|
||||||
// true the Machine is unrecoverable and the caller must abandon it.
|
|
||||||
type Machine struct {
|
|
||||||
hs *noise.HandshakeState
|
|
||||||
getCred GetCredentialFunc
|
|
||||||
allocIndex IndexAllocator
|
|
||||||
verifier CertVerifier
|
|
||||||
result *Result
|
|
||||||
msgs []msgFlags
|
|
||||||
myVersion cert.Version
|
|
||||||
subtype header.MessageSubType
|
|
||||||
indexAllocated bool
|
|
||||||
remoteCertSet bool
|
|
||||||
payloadSet bool
|
|
||||||
failed bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewMachine creates a handshake state machine. The subtype determines both
|
|
||||||
// the noise pattern and the per-message content layout. The credential for
|
|
||||||
// `version` is fetched via getCred and used to seed the noise.HandshakeState.
|
|
||||||
// IndexAllocator is called lazily when the first outgoing payload is built.
|
|
||||||
func NewMachine(
|
|
||||||
version cert.Version,
|
|
||||||
getCred GetCredentialFunc,
|
|
||||||
verifier CertVerifier,
|
|
||||||
allocIndex IndexAllocator,
|
|
||||||
initiator bool,
|
|
||||||
subtype header.MessageSubType,
|
|
||||||
) (*Machine, error) {
|
|
||||||
info, err := subtypeInfoFor(subtype)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
cred := getCred(version)
|
|
||||||
if cred == nil {
|
|
||||||
return nil, fmt.Errorf("%w: %v", ErrNoCredential, version)
|
|
||||||
}
|
|
||||||
|
|
||||||
hs, err := cred.buildHandshakeState(initiator, info.pattern)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("build noise state: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return &Machine{
|
|
||||||
hs: hs,
|
|
||||||
subtype: subtype,
|
|
||||||
msgs: info.msgs,
|
|
||||||
getCred: getCred,
|
|
||||||
allocIndex: allocIndex,
|
|
||||||
verifier: verifier,
|
|
||||||
myVersion: version,
|
|
||||||
result: &Result{
|
|
||||||
Initiator: initiator,
|
|
||||||
Cipher: cred.cipherSuite,
|
|
||||||
},
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Failed returns true if the Machine is in an unrecoverable state.
|
|
||||||
func (m *Machine) Failed() bool {
|
|
||||||
return m.failed
|
|
||||||
}
|
|
||||||
|
|
||||||
// Subtype returns the handshake subtype this Machine was built for.
|
|
||||||
func (m *Machine) Subtype() header.MessageSubType {
|
|
||||||
return m.subtype
|
|
||||||
}
|
|
||||||
|
|
||||||
// MessageIndex returns the noise handshake message index, which equals the
|
|
||||||
// wire counter of the most recently sent or received message.
|
|
||||||
func (m *Machine) MessageIndex() int {
|
|
||||||
return m.hs.MessageIndex()
|
|
||||||
}
|
|
||||||
|
|
||||||
// requireComplete checks that both a peer cert and payload have been received.
|
|
||||||
// Marks the machine as failed if not.
|
|
||||||
func (m *Machine) requireComplete() error {
|
|
||||||
if !m.payloadSet || !m.remoteCertSet {
|
|
||||||
m.failed = true
|
|
||||||
return ErrIncompleteHandshake
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// myMsgFlags returns the flags for the current outgoing message.
|
|
||||||
func (m *Machine) myMsgFlags() msgFlags {
|
|
||||||
idx := m.hs.MessageIndex()
|
|
||||||
if idx < len(m.msgs) {
|
|
||||||
return m.msgs[idx]
|
|
||||||
}
|
|
||||||
return msgFlags{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// peerMsgFlags returns the flags for the message we just read.
|
|
||||||
func (m *Machine) peerMsgFlags() msgFlags {
|
|
||||||
idx := m.hs.MessageIndex() - 1
|
|
||||||
if idx >= 0 && idx < len(m.msgs) {
|
|
||||||
return m.msgs[idx]
|
|
||||||
}
|
|
||||||
return msgFlags{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Initiate produces the first handshake message. Only valid for initiators,
|
|
||||||
// and must be called exactly once before ProcessPacket.
|
|
||||||
//
|
|
||||||
// out is a destination buffer the message is appended to and returned. Pass
|
|
||||||
// nil to allocate fresh, or pass a re-used buffer sliced to length 0 (e.g.
|
|
||||||
// buf[:0]) with sufficient capacity to avoid allocation.
|
|
||||||
//
|
|
||||||
// An error return may not indicate a fatal condition, check Failed() to
|
|
||||||
// determine if the Machine can still be used.
|
|
||||||
func (m *Machine) Initiate(out []byte) ([]byte, error) {
|
|
||||||
if m.failed {
|
|
||||||
return nil, ErrMachineFailed
|
|
||||||
}
|
|
||||||
if !m.result.Initiator {
|
|
||||||
m.failed = true
|
|
||||||
return nil, ErrInitiateOnResponder
|
|
||||||
}
|
|
||||||
if m.hs.MessageIndex() != 0 {
|
|
||||||
m.failed = true
|
|
||||||
return nil, ErrInitiateAlreadyCalled
|
|
||||||
}
|
|
||||||
|
|
||||||
// At MessageIndex=0 with RemoteIndex still zero, buildResponse produces
|
|
||||||
// header counter 1 and remote index 0, which is what the initial message needs.
|
|
||||||
out, _, _, err := m.buildResponse(out)
|
|
||||||
if err != nil {
|
|
||||||
m.failed = true
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// ProcessPacket handles an incoming handshake message. It advances the Noise
|
|
||||||
// state, validates the peer certificate via the verifier, and optionally
|
|
||||||
// produces a response.
|
|
||||||
//
|
|
||||||
// out is a destination buffer the response is appended to and returned. Pass
|
|
||||||
// nil to allocate fresh, or pass a re-used buffer sliced to length 0 (e.g.
|
|
||||||
// buf[:0]) with sufficient capacity to avoid allocation. The returned slice
|
|
||||||
// is nil when no outgoing message is produced (handshake complete on this
|
|
||||||
// side, or final message of a multi-message pattern).
|
|
||||||
//
|
|
||||||
// Returns a non-nil Result when the handshake is complete.
|
|
||||||
// An error return may not indicate a fatal condition, check Failed() to
|
|
||||||
// determine if the Machine can still be used.
|
|
||||||
func (m *Machine) ProcessPacket(out, packet []byte) ([]byte, *Result, error) {
|
|
||||||
if m.failed {
|
|
||||||
return nil, nil, ErrMachineFailed
|
|
||||||
}
|
|
||||||
if len(packet) < header.Len {
|
|
||||||
return nil, nil, ErrPacketTooShort
|
|
||||||
}
|
|
||||||
// Reject packets whose subtype doesn't match the one this Machine was
|
|
||||||
// built for. A pending handshake that suddenly receives a different
|
|
||||||
// subtype on its index is either a stray packet that matched by chance
|
|
||||||
// or a peer protocol violation; drop it without failing the Machine so
|
|
||||||
// the legitimate retransmit can still complete.
|
|
||||||
if header.MessageSubType(packet[1]) != m.subtype {
|
|
||||||
return nil, nil, ErrSubtypeMismatch
|
|
||||||
}
|
|
||||||
if m.result.Initiator && m.hs.MessageIndex() == 0 {
|
|
||||||
m.failed = true
|
|
||||||
return nil, nil, ErrInitiateNotCalled
|
|
||||||
}
|
|
||||||
|
|
||||||
// The (eKey, dKey) ordering here is correct for IX, where the initiator
|
|
||||||
// completes the handshake by reading the responder's stage-2 message.
|
|
||||||
// noise returns (cs1, cs2) where cs1 is the initiator->responder cipher.
|
|
||||||
// For 3-message patterns where a responder finishes by reading the final
|
|
||||||
// message, this ordering would be wrong; revisit when XX/pqIX lands.
|
|
||||||
msg, eKey, dKey, err := m.hs.ReadMessage(nil, packet[header.Len:])
|
|
||||||
if err != nil {
|
|
||||||
// Noise ReadMessage failed. The noise library checkpoints and rolls back
|
|
||||||
// on failure, so the Machine is still alive. The caller can retry with
|
|
||||||
// a different packet.
|
|
||||||
return nil, nil, fmt.Errorf("noise ReadMessage: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// From here on, noise state has advanced. Any error is fatal.
|
|
||||||
flags := m.peerMsgFlags()
|
|
||||||
|
|
||||||
if err := m.processPayload(msg, flags); err != nil {
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// If ReadMessage derived keys, the handshake is complete. Noise should
|
|
||||||
// always produce both keys together; asymmetry is a protocol invariant
|
|
||||||
// violation.
|
|
||||||
if eKey != nil || dKey != nil {
|
|
||||||
if eKey == nil || dKey == nil {
|
|
||||||
m.failed = true
|
|
||||||
return nil, nil, ErrAsymmetricCipherKeys
|
|
||||||
}
|
|
||||||
if err := m.requireComplete(); err != nil {
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
return nil, m.completed(eKey, dKey), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// ReadMessage didn't complete, produce the next outgoing message
|
|
||||||
out, dk, ek, err := m.buildResponse(out)
|
|
||||||
if err != nil {
|
|
||||||
m.failed = true
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if ek != nil || dk != nil {
|
|
||||||
if ek == nil || dk == nil {
|
|
||||||
m.failed = true
|
|
||||||
return nil, nil, ErrAsymmetricCipherKeys
|
|
||||||
}
|
|
||||||
if err := m.requireComplete(); err != nil {
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
return out, m.completed(ek, dk), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return out, nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Machine) completed(eKey, dKey *noise.CipherState) *Result {
|
|
||||||
m.result.EKey = eKey
|
|
||||||
m.result.DKey = dKey
|
|
||||||
m.result.MessageIndex = uint64(m.hs.MessageIndex())
|
|
||||||
return m.result
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Machine) processPayload(msg []byte, flags msgFlags) error {
|
|
||||||
if len(msg) == 0 {
|
|
||||||
if flags.expectsPayload || flags.expectsCert {
|
|
||||||
m.failed = true
|
|
||||||
return ErrMissingContent
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
payload, err := UnmarshalPayload(msg)
|
|
||||||
if err != nil {
|
|
||||||
m.failed = true
|
|
||||||
return fmt.Errorf("unmarshal handshake: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Assert the payload contains exactly what we expect
|
|
||||||
hasPayloadData := payload.InitiatorIndex != 0 || payload.ResponderIndex != 0 || payload.Time != 0
|
|
||||||
if hasPayloadData != flags.expectsPayload {
|
|
||||||
m.failed = true
|
|
||||||
return ErrUnexpectedContent
|
|
||||||
}
|
|
||||||
|
|
||||||
hasCertData := len(payload.Cert) > 0
|
|
||||||
if hasCertData != flags.expectsCert {
|
|
||||||
m.failed = true
|
|
||||||
return ErrUnexpectedContent
|
|
||||||
}
|
|
||||||
|
|
||||||
// Process payload
|
|
||||||
if flags.expectsPayload {
|
|
||||||
var remoteIndex uint32
|
|
||||||
if m.result.Initiator {
|
|
||||||
remoteIndex = payload.ResponderIndex
|
|
||||||
} else {
|
|
||||||
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.payloadSet = true
|
|
||||||
}
|
|
||||||
|
|
||||||
// Process certificate
|
|
||||||
if flags.expectsCert {
|
|
||||||
if err := m.validateCert(payload); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Machine) validateCert(payload Payload) error {
|
|
||||||
cred := m.getCred(m.myVersion)
|
|
||||||
if cred == nil {
|
|
||||||
m.failed = true
|
|
||||||
return fmt.Errorf("%w: %v", ErrNoCredential, m.myVersion)
|
|
||||||
}
|
|
||||||
rc, err := cert.Recombine(
|
|
||||||
cert.Version(payload.CertVersion),
|
|
||||||
payload.Cert,
|
|
||||||
m.hs.PeerStatic(),
|
|
||||||
cred.Cert.Curve(),
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
m.failed = true
|
|
||||||
return fmt.Errorf("recombine cert: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !bytes.Equal(rc.PublicKey(), m.hs.PeerStatic()) {
|
|
||||||
m.failed = true
|
|
||||||
return ErrPublicKeyMismatch
|
|
||||||
}
|
|
||||||
|
|
||||||
// Version negotiation, if the peer sent a different version and we have it, switch
|
|
||||||
if rc.Version() != m.myVersion {
|
|
||||||
if m.getCred(rc.Version()) != nil {
|
|
||||||
m.myVersion = rc.Version()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
verified, err := m.verifier(rc)
|
|
||||||
if err != nil {
|
|
||||||
m.failed = true
|
|
||||||
return fmt.Errorf("verify cert: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.result.RemoteCert = verified
|
|
||||||
m.remoteCertSet = true
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Machine) marshalOutgoing(flags msgFlags) ([]byte, error) {
|
|
||||||
if !flags.expectsPayload && !flags.expectsCert {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var p Payload
|
|
||||||
if flags.expectsPayload {
|
|
||||||
if !m.indexAllocated {
|
|
||||||
index, err := m.allocIndex()
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("%w: %w", ErrIndexAllocation, err)
|
|
||||||
}
|
|
||||||
m.result.LocalIndex = index
|
|
||||||
m.indexAllocated = true
|
|
||||||
}
|
|
||||||
|
|
||||||
if m.result.Initiator {
|
|
||||||
p.InitiatorIndex = m.result.LocalIndex
|
|
||||||
} else {
|
|
||||||
p.ResponderIndex = m.result.LocalIndex
|
|
||||||
p.InitiatorIndex = m.result.RemoteIndex
|
|
||||||
}
|
|
||||||
p.Time = uint64(time.Now().UnixNano())
|
|
||||||
}
|
|
||||||
if flags.expectsCert {
|
|
||||||
cred := m.getCred(m.myVersion)
|
|
||||||
if cred == nil {
|
|
||||||
return nil, fmt.Errorf("%w: %v", ErrNoCredential, m.myVersion)
|
|
||||||
}
|
|
||||||
p.Cert = cred.Bytes
|
|
||||||
p.CertVersion = uint32(cred.Cert.Version())
|
|
||||||
m.result.MyCert = cred.Cert
|
|
||||||
}
|
|
||||||
|
|
||||||
return MarshalPayload(nil, p), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Machine) buildResponse(out []byte) ([]byte, *noise.CipherState, *noise.CipherState, error) {
|
|
||||||
flags := m.myMsgFlags()
|
|
||||||
hsBytes, err := m.marshalOutgoing(flags)
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Extend out by header.Len to make room for the header. slices.Grow is a
|
|
||||||
// no-op when the cap is already sufficient (the zero-copy case where the
|
|
||||||
// caller passed a pre-sized buffer). header.Encode overwrites the new
|
|
||||||
// bytes, so they don't need to be zeroed.
|
|
||||||
start := len(out)
|
|
||||||
out = slices.Grow(out, header.Len)[:start+header.Len]
|
|
||||||
header.Encode(
|
|
||||||
out[start:],
|
|
||||||
header.Version, header.Handshake, m.subtype,
|
|
||||||
m.result.RemoteIndex,
|
|
||||||
uint64(m.hs.MessageIndex()+1),
|
|
||||||
)
|
|
||||||
|
|
||||||
// noise.WriteMessage appends the encrypted handshake message to out,
|
|
||||||
// reusing capacity when present.
|
|
||||||
//
|
|
||||||
// The (dKey, eKey) ordering here is correct for IX, where the responder
|
|
||||||
// completes the handshake by writing the stage-2 message. noise returns
|
|
||||||
// (cs1, cs2) where cs1 is the initiator->responder cipher (which is the
|
|
||||||
// responder's decrypt key). For 3-message patterns where an initiator
|
|
||||||
// finishes by writing the final message, this ordering would be wrong;
|
|
||||||
// revisit when XX/pqIX lands.
|
|
||||||
out, dKey, eKey, err := m.hs.WriteMessage(out, hsBytes)
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, nil, fmt.Errorf("noise WriteMessage: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return out, dKey, eKey, nil
|
|
||||||
}
|
|
||||||
@@ -1,680 +0,0 @@
|
|||||||
package handshake
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
ct "github.com/slackhq/nebula/cert_test"
|
|
||||||
"github.com/slackhq/nebula/header"
|
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestMachineIXHappyPath(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
|
|
||||||
initCS := newTestCertState(t, ca, caKey, "initiator", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
respCS := newTestCertState(t, ca, caKey, "responder", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
|
||||||
|
|
||||||
initR, respR := doFullHandshake(t, initCS, respCS, caPool)
|
|
||||||
|
|
||||||
assert.Equal(t, "responder", initR.RemoteCert.Certificate.Name())
|
|
||||||
assert.Equal(t, "initiator", respR.RemoteCert.Certificate.Name())
|
|
||||||
|
|
||||||
assert.Equal(t, uint32(1000), initR.LocalIndex)
|
|
||||||
assert.Equal(t, uint32(2000), initR.RemoteIndex)
|
|
||||||
assert.Equal(t, uint32(2000), respR.LocalIndex)
|
|
||||||
assert.Equal(t, uint32(1000), respR.RemoteIndex)
|
|
||||||
|
|
||||||
assert.Equal(t, uint64(2), initR.MessageIndex, "IX has 2 messages")
|
|
||||||
assert.Equal(t, uint64(2), respR.MessageIndex, "IX has 2 messages")
|
|
||||||
|
|
||||||
ct1, err := initR.EKey.Encrypt(nil, nil, []byte("hello"))
|
|
||||||
require.NoError(t, err)
|
|
||||||
pt1, err := respR.DKey.Decrypt(nil, nil, ct1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, []byte("hello"), pt1)
|
|
||||||
|
|
||||||
ct2, err := respR.EKey.Encrypt(nil, nil, []byte("world"))
|
|
||||||
require.NoError(t, err)
|
|
||||||
pt2, err := initR.DKey.Decrypt(nil, nil, ct2)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, []byte("world"), pt2)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMachineInitiateErrors(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
v := testVerifier(caPool)
|
|
||||||
|
|
||||||
t.Run("initiate on responder", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
_, err := m.Initiate(nil)
|
|
||||||
require.ErrorIs(t, err, ErrInitiateOnResponder)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("initiate called twice", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, true, 100)
|
|
||||||
_, err := m.Initiate(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = m.Initiate(nil)
|
|
||||||
require.ErrorIs(t, err, ErrInitiateAlreadyCalled)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("process packet before initiate on initiator", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, true, 100)
|
|
||||||
_, _, err := m.ProcessPacket(nil, make([]byte, 100))
|
|
||||||
require.ErrorIs(t, err, ErrInitiateNotCalled)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("calling failed machine", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
_, err := m.Initiate(nil) // fails: responder
|
|
||||||
require.Error(t, err)
|
|
||||||
_, err = m.Initiate(nil) // fails: already failed
|
|
||||||
require.ErrorIs(t, err, ErrMachineFailed)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMachineProcessPacketErrors(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
v := testVerifier(caPool)
|
|
||||||
|
|
||||||
t.Run("packet too short", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
_, _, err := m.ProcessPacket(nil, []byte{1, 2, 3})
|
|
||||||
require.ErrorIs(t, err, ErrPacketTooShort)
|
|
||||||
assert.False(t, m.Failed(), "short packet should not kill machine")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("noise decryption failure is recoverable", func(t *testing.T) {
|
|
||||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
initM := newTestMachine(t, initCS, v, true, 100)
|
|
||||||
msg1, err := initM.Initiate(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
respM := newTestMachine(t, cs, v, false, 200)
|
|
||||||
resp, _, err := respM.ProcessPacket(nil, msg1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
corrupted := make([]byte, len(resp))
|
|
||||||
copy(corrupted, resp)
|
|
||||||
for i := header.Len; i < len(corrupted); i++ {
|
|
||||||
corrupted[i] ^= 0xff
|
|
||||||
}
|
|
||||||
_, _, err = initM.ProcessPacket(nil, corrupted)
|
|
||||||
require.Error(t, err)
|
|
||||||
assert.False(t, initM.Failed(), "noise failure should be recoverable")
|
|
||||||
|
|
||||||
// And the machine should still complete a real handshake afterward.
|
|
||||||
_, result, err := initM.ProcessPacket(nil, resp)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, result, "initiator should complete on the legitimate response")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("invalid cert is fatal", func(t *testing.T) {
|
|
||||||
otherCA, _, otherCAKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
otherCS := newTestCertState(t, otherCA, otherCAKey, "other", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
|
||||||
|
|
||||||
initM := newTestMachine(t, otherCS, testVerifier(ct.NewTestCAPool(otherCA)), true, 100)
|
|
||||||
msg1, err := initM.Initiate(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
respM := newTestMachine(t, cs, v, false, 200)
|
|
||||||
_, _, err = respM.ProcessPacket(nil, msg1)
|
|
||||||
require.Error(t, err)
|
|
||||||
assert.True(t, respM.Failed(), "cert validation failure should kill machine")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("subtype mismatch is recoverable", func(t *testing.T) {
|
|
||||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
initM := newTestMachine(t, initCS, v, true, 100)
|
|
||||||
msg1, err := initM.Initiate(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Mutate the subtype byte (offset 1 in the header) to a value the
|
|
||||||
// responder Machine wasn't built for.
|
|
||||||
bad := make([]byte, len(msg1))
|
|
||||||
copy(bad, msg1)
|
|
||||||
bad[1] = 0xff
|
|
||||||
|
|
||||||
respM := newTestMachine(t, cs, v, false, 200)
|
|
||||||
_, _, err = respM.ProcessPacket(nil, bad)
|
|
||||||
require.ErrorIs(t, err, ErrSubtypeMismatch)
|
|
||||||
assert.False(t, respM.Failed(), "subtype mismatch should not kill the machine")
|
|
||||||
|
|
||||||
// And the machine should still complete a real handshake afterward.
|
|
||||||
resp, result, err := respM.ProcessPacket(nil, msg1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, result, "responder should complete on the legitimate stage-1 packet")
|
|
||||||
assert.NotEmpty(t, resp, "responder should produce a stage-2 reply")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMachineProcessPayload exercises processPayload's internal validation
|
|
||||||
// directly. Most of these failure modes can't be reached black-box once the
|
|
||||||
// subtype check at the top of ProcessPacket gates external callers, so we
|
|
||||||
// drive them by hand here for coverage.
|
|
||||||
func TestMachineProcessPayload(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
v := testVerifier(caPool)
|
|
||||||
|
|
||||||
t.Run("empty message with expects fails", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
err := m.processPayload(nil, msgFlags{expectsPayload: true, expectsCert: true})
|
|
||||||
require.ErrorIs(t, err, ErrMissingContent)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("empty message with no expects passes", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
err := m.processPayload(nil, msgFlags{})
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.False(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("malformed protobuf is fatal", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
err := m.processPayload([]byte{0xff, 0xff, 0xff}, msgFlags{expectsPayload: true, expectsCert: true})
|
|
||||||
require.Error(t, err)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("unexpected payload data is fatal", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
// A payload with index data when none was expected.
|
|
||||||
bytes := MarshalPayload(nil, Payload{InitiatorIndex: 42, Time: 1})
|
|
||||||
err := m.processPayload(bytes, msgFlags{expectsPayload: false, expectsCert: false})
|
|
||||||
require.ErrorIs(t, err, ErrUnexpectedContent)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("unexpected cert data is fatal", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
// A payload with cert when none was expected.
|
|
||||||
bytes := MarshalPayload(nil, Payload{Cert: []byte{1, 2, 3}, CertVersion: 2})
|
|
||||||
err := m.processPayload(bytes, msgFlags{expectsPayload: false, expectsCert: false})
|
|
||||||
require.ErrorIs(t, err, ErrUnexpectedContent)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("missing payload data when expected is fatal", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
// Cert present, but no index/time fields.
|
|
||||||
bytes := MarshalPayload(nil, Payload{Cert: []byte{1, 2, 3}, CertVersion: 2})
|
|
||||||
err := m.processPayload(bytes, msgFlags{expectsPayload: true, expectsCert: true})
|
|
||||||
require.ErrorIs(t, err, ErrUnexpectedContent)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
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
|
|
||||||
// directly. Like processPayload above this isn't reachable from a normal IX
|
|
||||||
// flow, so we drive it by hand.
|
|
||||||
func TestMachineRequireComplete(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
v := testVerifier(caPool)
|
|
||||||
|
|
||||||
t.Run("missing both fails", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
err := m.requireComplete()
|
|
||||||
require.ErrorIs(t, err, ErrIncompleteHandshake)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("payload only fails", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
m.payloadSet = true
|
|
||||||
err := m.requireComplete()
|
|
||||||
require.ErrorIs(t, err, ErrIncompleteHandshake)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("cert only fails", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
m.remoteCertSet = true
|
|
||||||
err := m.requireComplete()
|
|
||||||
require.ErrorIs(t, err, ErrIncompleteHandshake)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("both set passes", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
m.payloadSet = true
|
|
||||||
m.remoteCertSet = true
|
|
||||||
err := m.requireComplete()
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.False(t, m.Failed())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMachineAESCipher(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
|
|
||||||
initCS := newTestCertStateWithCipher(
|
|
||||||
t, ca, caKey, "init",
|
|
||||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
|
||||||
noiseutil.CipherAESGCM,
|
|
||||||
)
|
|
||||||
respCS := newTestCertStateWithCipher(
|
|
||||||
t, ca, caKey, "resp",
|
|
||||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
|
||||||
noiseutil.CipherAESGCM,
|
|
||||||
)
|
|
||||||
|
|
||||||
initR, respR := doFullHandshake(t, initCS, respCS, caPool)
|
|
||||||
|
|
||||||
ct1, err := initR.EKey.Encrypt(nil, nil, []byte("works"))
|
|
||||||
require.NoError(t, err)
|
|
||||||
pt1, err := respR.DKey.Decrypt(nil, nil, ct1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, []byte("works"), pt1)
|
|
||||||
|
|
||||||
ct2, err := respR.EKey.Encrypt(nil, nil, []byte("back"))
|
|
||||||
require.NoError(t, err)
|
|
||||||
pt2, err := initR.DKey.Decrypt(nil, nil, ct2)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, []byte("back"), pt2)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestResultFields(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
|
||||||
|
|
||||||
initR, respR := doFullHandshake(t, initCS, respCS, caPool)
|
|
||||||
|
|
||||||
assert.True(t, initR.Initiator)
|
|
||||||
assert.False(t, respR.Initiator)
|
|
||||||
assert.NotZero(t, initR.HandshakeTime)
|
|
||||||
assert.NotZero(t, respR.HandshakeTime)
|
|
||||||
assert.NotNil(t, initR.RemoteCert)
|
|
||||||
assert.NotNil(t, respR.RemoteCert)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMachineBufferReuse(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
|
||||||
v := testVerifier(caPool)
|
|
||||||
|
|
||||||
initM := newTestMachine(t, initCS, v, true, 1000)
|
|
||||||
respM := newTestMachine(t, respCS, v, false, 2000)
|
|
||||||
|
|
||||||
msg1, err := initM.Initiate(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
t.Run("response writes into provided buffer", func(t *testing.T) {
|
|
||||||
buf := make([]byte, 0, 4096)
|
|
||||||
resp, result, err := respM.ProcessPacket(buf, msg1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, result)
|
|
||||||
|
|
||||||
assert.NotEmpty(t, resp, "response should have content")
|
|
||||||
assert.Equal(t, &buf[:1][0], &resp[:1][0],
|
|
||||||
"response should reuse the provided buffer's backing array")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("initiate writes into provided buffer", func(t *testing.T) {
|
|
||||||
initM2 := newTestMachine(t, initCS, v, true, 3000)
|
|
||||||
buf := make([]byte, 0, 4096)
|
|
||||||
msg, err := initM2.Initiate(buf)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.NotEmpty(t, msg, "initiate should have content")
|
|
||||||
assert.Equal(t, &buf[:1][0], &msg[:1][0],
|
|
||||||
"initiate should reuse the provided buffer's backing array")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("nil out still works", func(t *testing.T) {
|
|
||||||
initM2 := newTestMachine(t, initCS, v, true, 4000)
|
|
||||||
respM2 := newTestMachine(t, respCS, v, false, 5000)
|
|
||||||
|
|
||||||
msg1, err := initM2.Initiate(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
resp, _, err := respM2.ProcessPacket(nil, msg1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
out, result, err := initM2.ProcessPacket(nil, resp)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.NotNil(t, result)
|
|
||||||
assert.Nil(t, out, "initiator should have no response for IX msg2")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMachineMsgIndexTracking(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
|
||||||
v := testVerifier(caPool)
|
|
||||||
|
|
||||||
initM := newTestMachine(t, initCS, v, true, 100)
|
|
||||||
respM := newTestMachine(t, respCS, v, false, 200)
|
|
||||||
|
|
||||||
msg1, err := initM.Initiate(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
resp1, result1, err := respM.ProcessPacket(nil, msg1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.NotNil(t, result1)
|
|
||||||
|
|
||||||
_, result2, err := initM.ProcessPacket(nil, resp1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.NotNil(t, result2)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMachineThreeMessagePattern(t *testing.T) {
|
|
||||||
registerTestXXInfo(t)
|
|
||||||
|
|
||||||
// Use HandshakeXX (3 messages) to verify the Machine handles multi-message
|
|
||||||
// patterns correctly. XX flow:
|
|
||||||
// msg1 (I->R): [E] - payload only, no cert
|
|
||||||
// msg2 (R->I): [E, ee, S, es] - payload + cert
|
|
||||||
// msg3 (I->R): [S, se] - cert only (no payload, not first two)
|
|
||||||
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
v := testVerifier(caPool)
|
|
||||||
|
|
||||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
|
||||||
|
|
||||||
initM, err := NewMachine(
|
|
||||||
cert.Version2,
|
|
||||||
initCS.getCredential, v,
|
|
||||||
func() (uint32, error) { return 1000, nil },
|
|
||||||
true, header.HandshakeXXPSK0,
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
respM, err := NewMachine(
|
|
||||||
cert.Version2,
|
|
||||||
respCS.getCredential, v,
|
|
||||||
func() (uint32, error) { return 2000, nil },
|
|
||||||
false, header.HandshakeXXPSK0,
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// msg1: initiator -> responder (E only, no cert)
|
|
||||||
msg1, err := initM.Initiate(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.NotEmpty(t, msg1)
|
|
||||||
|
|
||||||
// Responder processes msg1, should not complete yet, should produce msg2
|
|
||||||
msg2, result, err := respM.ProcessPacket(nil, msg1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Nil(t, result, "XX should not complete on msg1")
|
|
||||||
assert.NotEmpty(t, msg2, "responder should produce msg2")
|
|
||||||
|
|
||||||
// Initiator processes msg2: gets responder's cert, produces msg3, and
|
|
||||||
// completes (WriteMessage for msg3 derives keys)
|
|
||||||
msg3, initResult, err := initM.ProcessPacket(nil, msg2)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, initResult, "XX initiator should complete after reading msg2 and writing msg3")
|
|
||||||
assert.NotEmpty(t, msg3, "initiator should produce msg3")
|
|
||||||
assert.Equal(t, "resp", initResult.RemoteCert.Certificate.Name())
|
|
||||||
|
|
||||||
// Responder processes msg3: gets initiator's cert and completes
|
|
||||||
_, respResult, err := respM.ProcessPacket(nil, msg3)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, respResult, "XX responder should complete on msg3")
|
|
||||||
assert.Equal(t, "init", respResult.RemoteCert.Certificate.Name())
|
|
||||||
|
|
||||||
assert.Equal(t, uint64(3), initResult.MessageIndex, "XX has 3 messages")
|
|
||||||
assert.Equal(t, uint64(3), respResult.MessageIndex, "XX has 3 messages")
|
|
||||||
|
|
||||||
// Verify keys work
|
|
||||||
ct1, err := initResult.EKey.Encrypt(nil, nil, []byte("three messages"))
|
|
||||||
require.NoError(t, err)
|
|
||||||
pt1, err := respResult.DKey.Decrypt(nil, nil, ct1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, []byte("three messages"), pt1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// NOTE: ErrIncompleteHandshake is tested implicitly. It can't be triggered with
|
|
||||||
// IX since the cert is always in the payload. A 3-message pattern test (HybridIX)
|
|
||||||
// should exercise the case where cert arrives in msg3 and verify that completing
|
|
||||||
// without it fails.
|
|
||||||
|
|
||||||
func TestMachineExpiredCert(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519,
|
|
||||||
time.Now().Add(-24*time.Hour), time.Now().Add(24*time.Hour),
|
|
||||||
nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
|
|
||||||
expCert, _, expKeyPEM, _ := ct.NewTestCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, ca, caKey,
|
|
||||||
"expired", time.Now().Add(-2*time.Hour), time.Now().Add(-1*time.Hour),
|
|
||||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}, nil, nil,
|
|
||||||
)
|
|
||||||
expKey, _, _, err := cert.UnmarshalPrivateKeyFromPEM(expKeyPEM)
|
|
||||||
require.NoError(t, err)
|
|
||||||
expHsBytes, err := expCert.MarshalForHandshakes()
|
|
||||||
require.NoError(t, err)
|
|
||||||
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
|
||||||
|
|
||||||
expiredCS := &testCertState{
|
|
||||||
version: cert.Version2,
|
|
||||||
creds: map[cert.Version]*Credential{
|
|
||||||
cert.Version2: NewCredential(expCert, expHsBytes, expKey, ncs),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
respCS := newTestCertState(
|
|
||||||
t, ca, caKey, "responder",
|
|
||||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
|
||||||
)
|
|
||||||
|
|
||||||
_, respM, _, _, err := initiateHandshake(
|
|
||||||
t, expiredCS, testVerifier(caPool),
|
|
||||||
respCS, testVerifier(caPool),
|
|
||||||
)
|
|
||||||
require.ErrorContains(t, err, "verify cert")
|
|
||||||
assert.True(t, respM.Failed())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMachineNoCertNetworks(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
|
|
||||||
caHsBytes, err := ca.MarshalForHandshakes()
|
|
||||||
require.NoError(t, err)
|
|
||||||
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
|
||||||
|
|
||||||
noNetCS := &testCertState{
|
|
||||||
version: cert.Version2,
|
|
||||||
creds: map[cert.Version]*Credential{
|
|
||||||
cert.Version2: NewCredential(ca, caHsBytes, caKey, ncs),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
respCS := newTestCertState(
|
|
||||||
t, ca, caKey, "responder",
|
|
||||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
|
||||||
)
|
|
||||||
|
|
||||||
_, respM, _, _, err := initiateHandshake(
|
|
||||||
t, noNetCS, testVerifier(caPool),
|
|
||||||
respCS, testVerifier(caPool),
|
|
||||||
)
|
|
||||||
require.Error(t, err)
|
|
||||||
assert.True(t, respM.Failed())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMachineDifferentCAs(t *testing.T) {
|
|
||||||
ca1, _, caKey1, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
ca2, _, caKey2, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
|
|
||||||
initCS := newTestCertState(
|
|
||||||
t, ca1, caKey1, "init",
|
|
||||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
|
||||||
)
|
|
||||||
respCS := newTestCertState(
|
|
||||||
t, ca2, caKey2, "resp",
|
|
||||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
|
||||||
)
|
|
||||||
|
|
||||||
_, respM, _, _, err := initiateHandshake(
|
|
||||||
t, initCS, testVerifier(ct.NewTestCAPool(ca1)),
|
|
||||||
respCS, testVerifier(ct.NewTestCAPool(ca2)),
|
|
||||||
)
|
|
||||||
require.ErrorContains(t, err, "verify cert")
|
|
||||||
assert.True(t, respM.Failed())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMachineVersionNegotiation(t *testing.T) {
|
|
||||||
ca1, _, caKey1, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version1, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
ca2, _, caKey2, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca1, ca2)
|
|
||||||
|
|
||||||
makeMultiVersionResp := func(t *testing.T) *testCertState {
|
|
||||||
t.Helper()
|
|
||||||
respCertV1, _, respKeyPEM, _ := ct.NewTestCert(
|
|
||||||
cert.Version1, cert.Curve_CURVE25519, ca1, caKey1, "resp",
|
|
||||||
ca1.NotBefore(), ca1.NotAfter(),
|
|
||||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")}, nil, nil,
|
|
||||||
)
|
|
||||||
respKey, _, _, _ := cert.UnmarshalPrivateKeyFromPEM(respKeyPEM)
|
|
||||||
respCertV2, _ := ct.NewTestCertDifferentVersion(respCertV1, cert.Version2, ca2, caKey2)
|
|
||||||
respHsV1, _ := respCertV1.MarshalForHandshakes()
|
|
||||||
respHsV2, _ := respCertV2.MarshalForHandshakes()
|
|
||||||
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
|
||||||
return &testCertState{
|
|
||||||
version: cert.Version1,
|
|
||||||
creds: map[cert.Version]*Credential{
|
|
||||||
cert.Version1: NewCredential(respCertV1, respHsV1, respKey, ncs),
|
|
||||||
cert.Version2: NewCredential(respCertV2, respHsV2, respKey, ncs),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Run("responder matches initiator version", func(t *testing.T) {
|
|
||||||
initCS := newTestCertState(
|
|
||||||
t, ca2, caKey2, "init",
|
|
||||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
|
||||||
)
|
|
||||||
respCS := makeMultiVersionResp(t)
|
|
||||||
v := testVerifier(caPool)
|
|
||||||
|
|
||||||
initM, _, respResult, resp, err := initiateHandshake(
|
|
||||||
t, initCS, v,
|
|
||||||
respCS, v,
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, respResult)
|
|
||||||
|
|
||||||
assert.Equal(t, cert.Version2, respResult.MyCert.Version(),
|
|
||||||
"responder should negotiate to initiator's version")
|
|
||||||
|
|
||||||
_, initResult, err := initM.ProcessPacket(nil, resp)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, initResult)
|
|
||||||
assert.Equal(t, cert.Version2, initResult.RemoteCert.Certificate.Version(),
|
|
||||||
"initiator should see V2 cert from responder")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("responder keeps version when no match available", func(t *testing.T) {
|
|
||||||
initCS := newTestCertState(
|
|
||||||
t, ca2, caKey2, "init",
|
|
||||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
|
||||||
)
|
|
||||||
|
|
||||||
respCert, _, respKeyPEM, _ := ct.NewTestCert(
|
|
||||||
cert.Version1, cert.Curve_CURVE25519, ca1, caKey1, "resp",
|
|
||||||
ca1.NotBefore(), ca1.NotAfter(),
|
|
||||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")}, nil, nil,
|
|
||||||
)
|
|
||||||
respKey, _, _, _ := cert.UnmarshalPrivateKeyFromPEM(respKeyPEM)
|
|
||||||
respHs, _ := respCert.MarshalForHandshakes()
|
|
||||||
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
|
||||||
respCS := &testCertState{
|
|
||||||
version: cert.Version1,
|
|
||||||
creds: map[cert.Version]*Credential{
|
|
||||||
cert.Version1: NewCredential(respCert, respHs, respKey, ncs),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
v := testVerifier(caPool)
|
|
||||||
_, _, respResult, _, err := initiateHandshake(
|
|
||||||
t, initCS, v,
|
|
||||||
respCS, v,
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, respResult)
|
|
||||||
|
|
||||||
assert.Equal(t, cert.Version1, respResult.MyCert.Version(),
|
|
||||||
"responder should keep V1 when V2 not available")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@@ -1,54 +0,0 @@
|
|||||||
package handshake
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
"github.com/slackhq/nebula/header"
|
|
||||||
)
|
|
||||||
|
|
||||||
// msgFlags tracks what application data a handshake message carries.
|
|
||||||
type msgFlags struct {
|
|
||||||
expectsPayload bool // message carries indexes and time
|
|
||||||
expectsCert bool // message carries the certificate
|
|
||||||
}
|
|
||||||
|
|
||||||
// subtypeInfo bundles the noise pattern with the per-message flags for a
|
|
||||||
// given handshake subtype.
|
|
||||||
type subtypeInfo struct {
|
|
||||||
pattern noise.HandshakePattern
|
|
||||||
msgs []msgFlags
|
|
||||||
}
|
|
||||||
|
|
||||||
// subtypeInfos defines the noise pattern and message content layout for each
|
|
||||||
// handshake subtype.
|
|
||||||
var subtypeInfos = map[header.MessageSubType]subtypeInfo{
|
|
||||||
// IX: 2 messages, both carry payload and cert
|
|
||||||
header.HandshakeIXPSK0: {
|
|
||||||
pattern: noise.HandshakeIX,
|
|
||||||
msgs: []msgFlags{
|
|
||||||
{expectsPayload: true, expectsCert: true},
|
|
||||||
{expectsPayload: true, expectsCert: true},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
|
|
||||||
// XX: 3 messages
|
|
||||||
// msg1 (I->R): payload only
|
|
||||||
// msg2 (R->I): payload + cert
|
|
||||||
// msg3 (I->R): cert only
|
|
||||||
//header.HandshakeXXPSK0: {
|
|
||||||
// pattern: noise.HandshakeXX,
|
|
||||||
// msgs: []msgFlags{
|
|
||||||
// {expectsPayload: true, expectsCert: false},
|
|
||||||
// {expectsPayload: true, expectsCert: true},
|
|
||||||
// {expectsPayload: false, expectsCert: true},
|
|
||||||
// },
|
|
||||||
//},
|
|
||||||
}
|
|
||||||
|
|
||||||
func subtypeInfoFor(subtype header.MessageSubType) (subtypeInfo, error) {
|
|
||||||
if info, ok := subtypeInfos[subtype]; ok {
|
|
||||||
return info, nil
|
|
||||||
}
|
|
||||||
return subtypeInfo{}, fmt.Errorf("%w: %d", ErrUnknownSubtype, subtype)
|
|
||||||
}
|
|
||||||
@@ -1,63 +0,0 @@
|
|||||||
package handshake
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
"github.com/slackhq/nebula/header"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestSubtypeInfo(t *testing.T) {
|
|
||||||
t.Run("IX", func(t *testing.T) {
|
|
||||||
info, err := subtypeInfoFor(header.HandshakeIXPSK0)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, noise.HandshakeIX.Name, info.pattern.Name)
|
|
||||||
require.Len(t, info.msgs, 2)
|
|
||||||
// msg1: payload + cert
|
|
||||||
assert.True(t, info.msgs[0].expectsPayload)
|
|
||||||
assert.True(t, info.msgs[0].expectsCert)
|
|
||||||
// msg2: payload + cert
|
|
||||||
assert.True(t, info.msgs[1].expectsPayload)
|
|
||||||
assert.True(t, info.msgs[1].expectsCert)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("XX", func(t *testing.T) {
|
|
||||||
registerTestXXInfo(t)
|
|
||||||
info, err := subtypeInfoFor(header.HandshakeXXPSK0)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, noise.HandshakeXX.Name, info.pattern.Name)
|
|
||||||
require.Len(t, info.msgs, 3)
|
|
||||||
// msg1: payload only
|
|
||||||
assert.True(t, info.msgs[0].expectsPayload)
|
|
||||||
assert.False(t, info.msgs[0].expectsCert)
|
|
||||||
// msg2: payload + cert
|
|
||||||
assert.True(t, info.msgs[1].expectsPayload)
|
|
||||||
assert.True(t, info.msgs[1].expectsCert)
|
|
||||||
// msg3: cert only
|
|
||||||
assert.False(t, info.msgs[2].expectsPayload)
|
|
||||||
assert.True(t, info.msgs[2].expectsCert)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("unknown subtype returns error", func(t *testing.T) {
|
|
||||||
_, err := subtypeInfoFor(99)
|
|
||||||
require.ErrorIs(t, err, ErrUnknownSubtype)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// registerTestXXInfo temporarily registers XX subtype info for testing.
|
|
||||||
func registerTestXXInfo(t *testing.T) {
|
|
||||||
t.Helper()
|
|
||||||
subtypeInfos[header.HandshakeXXPSK0] = subtypeInfo{
|
|
||||||
pattern: noise.HandshakeXX,
|
|
||||||
msgs: []msgFlags{
|
|
||||||
{expectsPayload: true, expectsCert: false},
|
|
||||||
{expectsPayload: true, expectsCert: true},
|
|
||||||
{expectsPayload: false, expectsCert: true},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
t.Cleanup(func() {
|
|
||||||
delete(subtypeInfos, header.HandshakeXXPSK0)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@@ -1,173 +0,0 @@
|
|||||||
package handshake
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"math"
|
|
||||||
|
|
||||||
"google.golang.org/protobuf/encoding/protowire"
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
errInvalidHandshakeMessage = errors.New("invalid handshake message")
|
|
||||||
errInvalidHandshakeDetails = errors.New("invalid handshake details")
|
|
||||||
)
|
|
||||||
|
|
||||||
// Payload represents the decoded fields of a handshake message.
|
|
||||||
// Wire format is protobuf-compatible with NebulaHandshake{Details: NebulaHandshakeDetails{...}}.
|
|
||||||
type Payload struct {
|
|
||||||
Cert []byte
|
|
||||||
InitiatorIndex uint32
|
|
||||||
ResponderIndex uint32
|
|
||||||
Time uint64
|
|
||||||
CertVersion uint32
|
|
||||||
}
|
|
||||||
|
|
||||||
// Proto field numbers for NebulaHandshakeDetails
|
|
||||||
const (
|
|
||||||
fieldCert = 1 // bytes
|
|
||||||
fieldInitiatorIndex = 2 // uint32
|
|
||||||
fieldResponderIndex = 3 // uint32
|
|
||||||
fieldTime = 5 // uint64
|
|
||||||
fieldCertVersion = 8 // uint32
|
|
||||||
)
|
|
||||||
|
|
||||||
// MarshalPayload encodes a handshake payload in protobuf wire format compatible
|
|
||||||
// with NebulaHandshake{Details: NebulaHandshakeDetails{...}}.
|
|
||||||
// Returns out (which may be nil), with the marshalled Payload appended to it.
|
|
||||||
func MarshalPayload(out []byte, p Payload) []byte {
|
|
||||||
var details []byte
|
|
||||||
|
|
||||||
if len(p.Cert) > 0 {
|
|
||||||
details = protowire.AppendTag(details, fieldCert, protowire.BytesType)
|
|
||||||
details = protowire.AppendBytes(details, p.Cert)
|
|
||||||
}
|
|
||||||
if p.InitiatorIndex != 0 {
|
|
||||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, uint64(p.InitiatorIndex))
|
|
||||||
}
|
|
||||||
if p.ResponderIndex != 0 {
|
|
||||||
details = protowire.AppendTag(details, fieldResponderIndex, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, uint64(p.ResponderIndex))
|
|
||||||
}
|
|
||||||
if p.Time != 0 {
|
|
||||||
details = protowire.AppendTag(details, fieldTime, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, p.Time)
|
|
||||||
}
|
|
||||||
if p.CertVersion != 0 {
|
|
||||||
details = protowire.AppendTag(details, fieldCertVersion, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, uint64(p.CertVersion))
|
|
||||||
}
|
|
||||||
|
|
||||||
out = protowire.AppendTag(out, 1, protowire.BytesType)
|
|
||||||
out = protowire.AppendBytes(out, details)
|
|
||||||
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
// UnmarshalPayload decodes a protobuf-encoded NebulaHandshake message.
|
|
||||||
func UnmarshalPayload(b []byte) (Payload, error) {
|
|
||||||
var p Payload
|
|
||||||
|
|
||||||
for len(b) > 0 {
|
|
||||||
num, typ, n := protowire.ConsumeTag(b)
|
|
||||||
if n < 0 {
|
|
||||||
return p, errInvalidHandshakeMessage
|
|
||||||
}
|
|
||||||
b = b[n:]
|
|
||||||
|
|
||||||
switch {
|
|
||||||
case num == 1 && typ == protowire.BytesType:
|
|
||||||
details, n := protowire.ConsumeBytes(b)
|
|
||||||
if n < 0 {
|
|
||||||
return p, errInvalidHandshakeMessage
|
|
||||||
}
|
|
||||||
b = b[n:]
|
|
||||||
if err := unmarshalPayloadDetails(&p, details); err != nil {
|
|
||||||
return p, err
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
n := protowire.ConsumeFieldValue(num, typ, b)
|
|
||||||
if n < 0 {
|
|
||||||
return p, errInvalidHandshakeMessage
|
|
||||||
}
|
|
||||||
b = b[n:]
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return p, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func unmarshalPayloadDetails(p *Payload, b []byte) error {
|
|
||||||
for len(b) > 0 {
|
|
||||||
num, typ, n := protowire.ConsumeTag(b)
|
|
||||||
if n < 0 {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
b = b[n:]
|
|
||||||
|
|
||||||
// For known field numbers, reject any non-matching wire type as a
|
|
||||||
// hard error rather than silently skipping. The caller will catch
|
|
||||||
// missing-field cases downstream, but a wire-type mismatch on a tag
|
|
||||||
// we know is a peer protocol violation worth flagging here.
|
|
||||||
// Repeated occurrences of a singular field follow proto3 last-wins.
|
|
||||||
switch num {
|
|
||||||
case fieldCert:
|
|
||||||
if typ != protowire.BytesType {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
v, n := protowire.ConsumeBytes(b)
|
|
||||||
if n < 0 {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
p.Cert = append([]byte(nil), v...)
|
|
||||||
b = b[n:]
|
|
||||||
case fieldInitiatorIndex:
|
|
||||||
if typ != protowire.VarintType {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
v, n := protowire.ConsumeVarint(b)
|
|
||||||
if n < 0 || v > math.MaxUint32 {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
p.InitiatorIndex = uint32(v)
|
|
||||||
b = b[n:]
|
|
||||||
case fieldResponderIndex:
|
|
||||||
if typ != protowire.VarintType {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
v, n := protowire.ConsumeVarint(b)
|
|
||||||
if n < 0 || v > math.MaxUint32 {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
p.ResponderIndex = uint32(v)
|
|
||||||
b = b[n:]
|
|
||||||
case fieldTime:
|
|
||||||
if typ != protowire.VarintType {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
v, n := protowire.ConsumeVarint(b)
|
|
||||||
if n < 0 {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
p.Time = v
|
|
||||||
b = b[n:]
|
|
||||||
case fieldCertVersion:
|
|
||||||
if typ != protowire.VarintType {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
v, n := protowire.ConsumeVarint(b)
|
|
||||||
if n < 0 || v > math.MaxUint32 {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
p.CertVersion = uint32(v)
|
|
||||||
b = b[n:]
|
|
||||||
default:
|
|
||||||
n := protowire.ConsumeFieldValue(num, typ, b)
|
|
||||||
if n < 0 {
|
|
||||||
return errInvalidHandshakeDetails
|
|
||||||
}
|
|
||||||
b = b[n:]
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,361 +0,0 @@
|
|||||||
package handshake
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"math"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"google.golang.org/protobuf/encoding/protowire"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestPayloadRoundTrip(t *testing.T) {
|
|
||||||
t.Run("all fields set", func(t *testing.T) {
|
|
||||||
data := MarshalPayload(nil, Payload{
|
|
||||||
Cert: []byte("test-cert-bytes"),
|
|
||||||
CertVersion: 2,
|
|
||||||
InitiatorIndex: 12345,
|
|
||||||
ResponderIndex: 67890,
|
|
||||||
Time: 1234567890,
|
|
||||||
})
|
|
||||||
|
|
||||||
got, err := UnmarshalPayload(data)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, []byte("test-cert-bytes"), got.Cert)
|
|
||||||
assert.Equal(t, uint32(12345), got.InitiatorIndex)
|
|
||||||
assert.Equal(t, uint32(67890), got.ResponderIndex)
|
|
||||||
assert.Equal(t, uint64(1234567890), got.Time)
|
|
||||||
assert.Equal(t, uint32(2), got.CertVersion)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("minimal fields", func(t *testing.T) {
|
|
||||||
data := MarshalPayload(nil, Payload{InitiatorIndex: 1})
|
|
||||||
|
|
||||||
got, err := UnmarshalPayload(data)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, uint32(1), got.InitiatorIndex)
|
|
||||||
assert.Equal(t, uint32(0), got.ResponderIndex)
|
|
||||||
assert.Equal(t, uint64(0), got.Time)
|
|
||||||
assert.Nil(t, got.Cert)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("empty payload", func(t *testing.T) {
|
|
||||||
data := MarshalPayload(nil, Payload{})
|
|
||||||
|
|
||||||
got, err := UnmarshalPayload(data)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, uint32(0), got.InitiatorIndex)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("large cert bytes", func(t *testing.T) {
|
|
||||||
bigCert := make([]byte, 4096)
|
|
||||||
for i := range bigCert {
|
|
||||||
bigCert[i] = byte(i % 256)
|
|
||||||
}
|
|
||||||
|
|
||||||
data := MarshalPayload(nil, Payload{
|
|
||||||
Cert: bigCert,
|
|
||||||
CertVersion: 2,
|
|
||||||
InitiatorIndex: 999,
|
|
||||||
})
|
|
||||||
|
|
||||||
got, err := UnmarshalPayload(data)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, bigCert, got.Cert)
|
|
||||||
assert.Equal(t, uint32(999), got.InitiatorIndex)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("append to existing buffer", func(t *testing.T) {
|
|
||||||
prefix := []byte("prefix")
|
|
||||||
data := MarshalPayload(prefix, Payload{InitiatorIndex: 42})
|
|
||||||
|
|
||||||
assert.Equal(t, []byte("prefix"), data[:6])
|
|
||||||
|
|
||||||
got, err := UnmarshalPayload(data[6:])
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, uint32(42), got.InitiatorIndex)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPayloadUnknownFields(t *testing.T) {
|
|
||||||
t.Run("unknown field in outer message is skipped", func(t *testing.T) {
|
|
||||||
// Marshal a normal payload then append an unknown field (field 99, varint)
|
|
||||||
data := MarshalPayload(nil, Payload{InitiatorIndex: 42})
|
|
||||||
data = protowire.AppendTag(data, 99, protowire.VarintType)
|
|
||||||
data = protowire.AppendVarint(data, 12345)
|
|
||||||
|
|
||||||
got, err := UnmarshalPayload(data)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, uint32(42), got.InitiatorIndex)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("unknown field in details is skipped", func(t *testing.T) {
|
|
||||||
// Build details with a known field + unknown field
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, 77)
|
|
||||||
// Unknown field 50, varint
|
|
||||||
details = protowire.AppendTag(details, 50, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, 9999)
|
|
||||||
// Another known field after the unknown one
|
|
||||||
details = protowire.AppendTag(details, fieldResponderIndex, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, 88)
|
|
||||||
|
|
||||||
// Wrap in outer message
|
|
||||||
var data []byte
|
|
||||||
data = protowire.AppendTag(data, 1, protowire.BytesType)
|
|
||||||
data = protowire.AppendBytes(data, details)
|
|
||||||
|
|
||||||
got, err := UnmarshalPayload(data)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, uint32(77), got.InitiatorIndex)
|
|
||||||
assert.Equal(t, uint32(88), got.ResponderIndex)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("reserved fields 6 and 7 are skipped", func(t *testing.T) {
|
|
||||||
// Fields 6 and 7 are reserved in the proto definition
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, 100)
|
|
||||||
details = protowire.AppendTag(details, 6, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, 1)
|
|
||||||
details = protowire.AppendTag(details, 7, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, 2)
|
|
||||||
|
|
||||||
var data []byte
|
|
||||||
data = protowire.AppendTag(data, 1, protowire.BytesType)
|
|
||||||
data = protowire.AppendBytes(data, details)
|
|
||||||
|
|
||||||
got, err := UnmarshalPayload(data)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, uint32(100), got.InitiatorIndex)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPayloadBytesConsumed(t *testing.T) {
|
|
||||||
t.Run("all bytes consumed on valid input", func(t *testing.T) {
|
|
||||||
original := Payload{
|
|
||||||
Cert: []byte("cert"),
|
|
||||||
CertVersion: 2,
|
|
||||||
InitiatorIndex: 100,
|
|
||||||
ResponderIndex: 200,
|
|
||||||
Time: 999,
|
|
||||||
}
|
|
||||||
data := MarshalPayload(nil, original)
|
|
||||||
|
|
||||||
got, err := UnmarshalPayload(data)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Re-marshal and compare — proves we consumed and reproduced all fields
|
|
||||||
remarshaled := MarshalPayload(nil, got)
|
|
||||||
assert.Equal(t, data, remarshaled)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// wrapDetails wraps raw detail bytes in the outer NebulaHandshake envelope
|
|
||||||
// so UnmarshalPayload can reach unmarshalPayloadDetails.
|
|
||||||
func wrapDetails(details []byte) []byte {
|
|
||||||
var out []byte
|
|
||||||
out = protowire.AppendTag(out, 1, protowire.BytesType)
|
|
||||||
out = protowire.AppendBytes(out, details)
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPayloadUnmarshalErrors(t *testing.T) {
|
|
||||||
t.Run("nil input", func(t *testing.T) {
|
|
||||||
got, err := UnmarshalPayload(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, uint32(0), got.InitiatorIndex)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("truncated outer tag", func(t *testing.T) {
|
|
||||||
_, err := UnmarshalPayload([]byte{0x80})
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("truncated outer details field", func(t *testing.T) {
|
|
||||||
_, err := UnmarshalPayload([]byte{0x0a, 0x64, 0x01, 0x02, 0x03, 0x04, 0x05})
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("truncated outer unknown field", func(t *testing.T) {
|
|
||||||
// Valid tag for unknown field 99 varint, but no value follows
|
|
||||||
var data []byte
|
|
||||||
data = protowire.AppendTag(data, 99, protowire.VarintType)
|
|
||||||
_, err := UnmarshalPayload(data)
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("truncated details tag", func(t *testing.T) {
|
|
||||||
_, err := UnmarshalPayload(wrapDetails([]byte{0x80}))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("truncated cert bytes", func(t *testing.T) {
|
|
||||||
// Field 1 (cert), bytes type, length 10 but only 2 bytes
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldCert, protowire.BytesType)
|
|
||||||
details = append(details, 0x0a, 0x01, 0x02) // length 10, only 2 bytes
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("truncated initiator index varint", func(t *testing.T) {
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
|
||||||
details = append(details, 0x80) // incomplete varint
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("truncated responder index varint", func(t *testing.T) {
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldResponderIndex, protowire.VarintType)
|
|
||||||
details = append(details, 0x80)
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("truncated time varint", func(t *testing.T) {
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldTime, protowire.VarintType)
|
|
||||||
details = append(details, 0x80)
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("truncated cert version varint", func(t *testing.T) {
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldCertVersion, protowire.VarintType)
|
|
||||||
details = append(details, 0x80)
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("truncated unknown field in details", func(t *testing.T) {
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, 50, protowire.VarintType)
|
|
||||||
details = append(details, 0x80) // incomplete varint
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("cert with wrong wire type rejected", func(t *testing.T) {
|
|
||||||
// fieldCert as Varint instead of Bytes.
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldCert, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, 42)
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("initiator index with wrong wire type rejected", func(t *testing.T) {
|
|
||||||
// fieldInitiatorIndex as Bytes instead of Varint.
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.BytesType)
|
|
||||||
details = protowire.AppendBytes(details, []byte{1, 2, 3})
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("time with wrong wire type rejected", func(t *testing.T) {
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldTime, protowire.BytesType)
|
|
||||||
details = protowire.AppendBytes(details, []byte{1, 2, 3})
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("cert version with wrong wire type rejected", func(t *testing.T) {
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldCertVersion, protowire.BytesType)
|
|
||||||
details = protowire.AppendBytes(details, []byte{1, 2, 3})
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("repeated singular field follows proto3 last-wins", func(t *testing.T) {
|
|
||||||
// Per proto3, multiple instances of a singular field are accepted and
|
|
||||||
// the last value wins. We keep this behavior so that peers using
|
|
||||||
// alternative encoders aren't rejected.
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, 1)
|
|
||||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, 42)
|
|
||||||
got, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, uint32(42), got.InitiatorIndex)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("initiator index varint overflow rejected", func(t *testing.T) {
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, math.MaxUint32+1)
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("cert version varint overflow rejected", func(t *testing.T) {
|
|
||||||
var details []byte
|
|
||||||
details = protowire.AppendTag(details, fieldCertVersion, protowire.VarintType)
|
|
||||||
details = protowire.AppendVarint(details, math.MaxUint32+1)
|
|
||||||
_, err := UnmarshalPayload(wrapDetails(details))
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
// FuzzPayload feeds arbitrary bytes through UnmarshalPayload to confirm it
|
|
||||||
// never panics, and for any input that parses cleanly, that re-marshal +
|
|
||||||
// re-parse is a fix-point. Inputs come from an authenticated peer (post-
|
|
||||||
// noise-decrypt), so the threat model is "valid peer behaving arbitrarily,"
|
|
||||||
// not "unauthenticated injection."
|
|
||||||
func FuzzPayload(f *testing.F) {
|
|
||||||
// Seed corpus with a handful of known-good shapes.
|
|
||||||
f.Add(MarshalPayload(nil, Payload{}))
|
|
||||||
f.Add(MarshalPayload(nil, Payload{Cert: []byte{1, 2, 3}, CertVersion: 2}))
|
|
||||||
f.Add(MarshalPayload(nil, Payload{InitiatorIndex: 42, Time: 1}))
|
|
||||||
f.Add(MarshalPayload(nil, Payload{
|
|
||||||
Cert: []byte("seed-cert"),
|
|
||||||
InitiatorIndex: 1,
|
|
||||||
ResponderIndex: 2,
|
|
||||||
Time: 3,
|
|
||||||
CertVersion: 2,
|
|
||||||
}))
|
|
||||||
f.Add([]byte{})
|
|
||||||
f.Add([]byte{0xff})
|
|
||||||
|
|
||||||
f.Fuzz(func(t *testing.T, data []byte) {
|
|
||||||
p1, err := UnmarshalPayload(data)
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// For any input that parses, re-marshaling and re-parsing must
|
|
||||||
// yield an equivalent Payload. This catches dispatch bugs (e.g.
|
|
||||||
// emitting a field on marshal that we don't accept on parse) and
|
|
||||||
// any non-idempotent parsing behavior.
|
|
||||||
b2 := MarshalPayload(nil, p1)
|
|
||||||
p2, err := UnmarshalPayload(b2)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("re-parse of self-marshaled payload failed: %v\nintermediate: %x\n", err, b2)
|
|
||||||
}
|
|
||||||
if !payloadsEqual(p1, p2) {
|
|
||||||
t.Fatalf("re-marshal not idempotent\nfirst: %+v\nsecond: %+v", p1, p2)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func payloadsEqual(a, b Payload) bool {
|
|
||||||
return bytes.Equal(a.Cert, b.Cert) &&
|
|
||||||
a.InitiatorIndex == b.InitiatorIndex &&
|
|
||||||
a.ResponderIndex == b.ResponderIndex &&
|
|
||||||
a.Time == b.Time &&
|
|
||||||
a.CertVersion == b.CertVersion
|
|
||||||
}
|
|
||||||
+678
@@ -0,0 +1,678 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"net/netip"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
|
)
|
||||||
|
|
||||||
|
// NOISE IX Handshakes
|
||||||
|
|
||||||
|
// This function constructs a handshake packet, but does not actually send it
|
||||||
|
// Sending is done by the handshake manager
|
||||||
|
func ixHandshakeStage0(f *Interface, hh *HandshakeHostInfo) bool {
|
||||||
|
err := f.handshakeManager.allocateIndex(hh)
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||||
|
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).Error("Failed to generate index")
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
cs := f.pki.getCertState()
|
||||||
|
v := cs.initiatingVersion
|
||||||
|
if hh.initiatingVersionOverride != cert.VersionPre1 {
|
||||||
|
v = hh.initiatingVersionOverride
|
||||||
|
} else if v < cert.Version2 {
|
||||||
|
// If we're connecting to a v6 address we should encourage use of a V2 cert
|
||||||
|
for _, a := range hh.hostinfo.vpnAddrs {
|
||||||
|
if a.Is6() {
|
||||||
|
v = cert.Version2
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
crt := cs.getCertificate(v)
|
||||||
|
if crt == nil {
|
||||||
|
f.l.WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||||
|
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).
|
||||||
|
WithField("certVersion", v).
|
||||||
|
Error("Unable to handshake with host because no certificate is available")
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
crtHs := cs.getHandshakeBytes(v)
|
||||||
|
if crtHs == nil {
|
||||||
|
f.l.WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||||
|
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).
|
||||||
|
WithField("certVersion", v).
|
||||||
|
Error("Unable to handshake with host because no certificate handshake bytes is available")
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
ci, err := NewConnectionState(f.l, cs, crt, true, noise.HandshakeIX)
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||||
|
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).
|
||||||
|
WithField("certVersion", v).
|
||||||
|
Error("Failed to create connection state")
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
hh.hostinfo.ConnectionState = ci
|
||||||
|
|
||||||
|
hs := &NebulaHandshake{
|
||||||
|
Details: &NebulaHandshakeDetails{
|
||||||
|
InitiatorIndex: hh.hostinfo.localIndexId,
|
||||||
|
Time: uint64(time.Now().UnixNano()),
|
||||||
|
Cert: crtHs,
|
||||||
|
CertVersion: uint32(v),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
hsBytes, err := hs.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||||
|
WithField("certVersion", v).
|
||||||
|
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).Error("Failed to marshal handshake message")
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
h := header.Encode(make([]byte, header.Len), header.Version, header.Handshake, header.HandshakeIXPSK0, 0, 1)
|
||||||
|
|
||||||
|
msg, _, _, err := ci.H.WriteMessage(h, hsBytes)
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||||
|
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).Error("Failed to call noise.WriteMessage")
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// We are sending handshake packet 1, so we don't expect to receive
|
||||||
|
// handshake packet 1 from the responder
|
||||||
|
ci.window.Update(f.l, 1)
|
||||||
|
|
||||||
|
hh.hostinfo.HandshakePacket[0] = msg
|
||||||
|
hh.ready = true
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func ixHandshakeStage1(f *Interface, via ViaSender, packet []byte, h *header.H) {
|
||||||
|
cs := f.pki.getCertState()
|
||||||
|
crt := cs.GetDefaultCertificate()
|
||||||
|
if crt == nil {
|
||||||
|
f.l.WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).
|
||||||
|
WithField("certVersion", cs.initiatingVersion).
|
||||||
|
Error("Unable to handshake with host because no certificate is available")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ci, err := NewConnectionState(f.l, cs, crt, false, noise.HandshakeIX)
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||||
|
Error("Failed to create connection state")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Mark packet 1 as seen so it doesn't show up as missed
|
||||||
|
ci.window.Update(f.l, 1)
|
||||||
|
|
||||||
|
msg, _, _, err := ci.H.ReadMessage(nil, packet[header.Len:])
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||||
|
Error("Failed to call noise.ReadMessage")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hs := &NebulaHandshake{}
|
||||||
|
err = hs.Unmarshal(msg)
|
||||||
|
if err != nil || hs.Details == nil {
|
||||||
|
f.l.WithError(err).WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||||
|
Error("Failed unmarshal handshake message")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
rc, err := cert.Recombine(cert.Version(hs.Details.CertVersion), hs.Details.Cert, ci.H.PeerStatic(), ci.Curve())
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||||
|
Info("Handshake did not contain a certificate")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
remoteCert, err := f.pki.GetCAPool().VerifyCertificate(time.Now(), rc)
|
||||||
|
if err != nil {
|
||||||
|
fp, fperr := rc.Fingerprint()
|
||||||
|
if fperr != nil {
|
||||||
|
fp = "<error generating certificate fingerprint>"
|
||||||
|
}
|
||||||
|
|
||||||
|
e := f.l.WithError(err).WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||||
|
WithField("certVpnNetworks", rc.Networks()).
|
||||||
|
WithField("certFingerprint", fp)
|
||||||
|
|
||||||
|
if f.l.Level >= logrus.DebugLevel {
|
||||||
|
e = e.WithField("cert", rc)
|
||||||
|
}
|
||||||
|
|
||||||
|
e.Info("Invalid certificate from host")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if !bytes.Equal(remoteCert.Certificate.PublicKey(), ci.H.PeerStatic()) {
|
||||||
|
f.l.WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||||
|
WithField("cert", remoteCert).Info("public key mismatch between certificate and handshake")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if remoteCert.Certificate.Version() != ci.myCert.Version() {
|
||||||
|
// We started off using the wrong certificate version, lets see if we can match the version that was sent to us
|
||||||
|
myCertOtherVersion := cs.getCertificate(remoteCert.Certificate.Version())
|
||||||
|
if myCertOtherVersion == nil {
|
||||||
|
if f.l.Level >= logrus.DebugLevel {
|
||||||
|
f.l.WithError(err).WithFields(m{
|
||||||
|
"from": via,
|
||||||
|
"handshake": m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
"cert": remoteCert,
|
||||||
|
}).Debug("Might be unable to handshake with host due to missing certificate version")
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Record the certificate we are actually using
|
||||||
|
ci.myCert = myCertOtherVersion
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(remoteCert.Certificate.Networks()) == 0 {
|
||||||
|
f.l.WithError(err).WithField("from", via).
|
||||||
|
WithField("cert", remoteCert).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||||
|
Info("No networks in certificate")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
certName := remoteCert.Certificate.Name()
|
||||||
|
certVersion := remoteCert.Certificate.Version()
|
||||||
|
fingerprint := remoteCert.Fingerprint
|
||||||
|
issuer := remoteCert.Certificate.Issuer()
|
||||||
|
vpnNetworks := remoteCert.Certificate.Networks()
|
||||||
|
|
||||||
|
anyVpnAddrsInCommon := false
|
||||||
|
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
||||||
|
for i, network := range vpnNetworks {
|
||||||
|
if f.myVpnAddrsTable.Contains(network.Addr()) {
|
||||||
|
f.l.WithField("vpnNetworks", vpnNetworks).WithField("from", via).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
WithField("issuer", issuer).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).Error("Refusing to handshake with myself")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
vpnAddrs[i] = network.Addr()
|
||||||
|
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
||||||
|
anyVpnAddrsInCommon = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !via.IsRelayed {
|
||||||
|
// We only want to apply the remote allow list for direct tunnels here
|
||||||
|
if !f.lightHouse.GetRemoteAllowList().AllowAll(vpnAddrs, via.UdpAddr.Addr()) {
|
||||||
|
f.l.WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||||
|
Debug("lighthouse.remote_allow_list denied incoming handshake")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
myIndex, err := generateIndex(f.l)
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
WithField("issuer", issuer).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).Error("Failed to generate index")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo := &HostInfo{
|
||||||
|
ConnectionState: ci,
|
||||||
|
localIndexId: myIndex,
|
||||||
|
remoteIndexId: hs.Details.InitiatorIndex,
|
||||||
|
vpnAddrs: vpnAddrs,
|
||||||
|
HandshakePacket: make(map[uint8][]byte, 0),
|
||||||
|
lastHandshakeTime: hs.Details.Time,
|
||||||
|
relayState: RelayState{
|
||||||
|
relays: nil,
|
||||||
|
relayForByAddr: map[netip.Addr]*Relay{},
|
||||||
|
relayForByIdx: map[uint32]*Relay{},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgRxL := f.l.WithFields(m{
|
||||||
|
"vpnAddrs": vpnAddrs,
|
||||||
|
"from": via,
|
||||||
|
"certName": certName,
|
||||||
|
"certVersion": certVersion,
|
||||||
|
"fingerprint": fingerprint,
|
||||||
|
"issuer": issuer,
|
||||||
|
"initiatorIndex": hs.Details.InitiatorIndex,
|
||||||
|
"responderIndex": hs.Details.ResponderIndex,
|
||||||
|
"remoteIndex": h.RemoteIndex,
|
||||||
|
"handshake": m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
})
|
||||||
|
|
||||||
|
if anyVpnAddrsInCommon {
|
||||||
|
msgRxL.Info("Handshake message received")
|
||||||
|
} else {
|
||||||
|
//todo warn if not lighthouse or relay?
|
||||||
|
msgRxL.Info("Handshake message received, but no vpnNetworks in common.")
|
||||||
|
}
|
||||||
|
|
||||||
|
hs.Details.ResponderIndex = myIndex
|
||||||
|
hs.Details.Cert = cs.getHandshakeBytes(ci.myCert.Version())
|
||||||
|
if hs.Details.Cert == nil {
|
||||||
|
msgRxL.WithField("myCertVersion", ci.myCert.Version()).
|
||||||
|
Error("Unable to handshake with host because no certificate handshake bytes is available")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hs.Details.CertVersion = uint32(ci.myCert.Version())
|
||||||
|
// Update the time in case their clock is way off from ours
|
||||||
|
hs.Details.Time = uint64(time.Now().UnixNano())
|
||||||
|
|
||||||
|
hsBytes, err := hs.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
WithField("issuer", issuer).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).Error("Failed to marshal handshake message")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
nh := header.Encode(make([]byte, header.Len), header.Version, header.Handshake, header.HandshakeIXPSK0, hs.Details.InitiatorIndex, 2)
|
||||||
|
msg, dKey, eKey, err := ci.H.WriteMessage(nh, hsBytes)
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
WithField("issuer", issuer).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).Error("Failed to call noise.WriteMessage")
|
||||||
|
return
|
||||||
|
} else if dKey == nil || eKey == nil {
|
||||||
|
f.l.WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
WithField("issuer", issuer).
|
||||||
|
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).Error("Noise did not arrive at a key")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo.HandshakePacket[0] = make([]byte, len(packet[header.Len:]))
|
||||||
|
copy(hostinfo.HandshakePacket[0], packet[header.Len:])
|
||||||
|
|
||||||
|
// Regardless of whether you are the sender or receiver, you should arrive here
|
||||||
|
// and complete standing up the connection.
|
||||||
|
hostinfo.HandshakePacket[2] = make([]byte, len(msg))
|
||||||
|
copy(hostinfo.HandshakePacket[2], msg)
|
||||||
|
|
||||||
|
// We are sending handshake packet 2, so we don't expect to receive
|
||||||
|
// handshake packet 2 from the initiator.
|
||||||
|
ci.window.Update(f.l, 2)
|
||||||
|
|
||||||
|
ci.peerCert = remoteCert
|
||||||
|
ci.dKey = NewNebulaCipherState(dKey)
|
||||||
|
ci.eKey = NewNebulaCipherState(eKey)
|
||||||
|
|
||||||
|
hostinfo.remotes = f.lightHouse.QueryCache(vpnAddrs)
|
||||||
|
if !via.IsRelayed {
|
||||||
|
hostinfo.SetRemote(via.UdpAddr)
|
||||||
|
}
|
||||||
|
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
||||||
|
|
||||||
|
existing, err := f.handshakeManager.CheckAndComplete(hostinfo, 0, f)
|
||||||
|
if err != nil {
|
||||||
|
switch err {
|
||||||
|
case ErrAlreadySeen:
|
||||||
|
// Update remote if preferred
|
||||||
|
if existing.SetRemoteIfPreferred(f.hostMap, via) {
|
||||||
|
// Send a test packet to ensure the other side has also switched to
|
||||||
|
// the preferred remote
|
||||||
|
f.SendMessageToVpnAddr(header.Test, header.TestRequest, vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
||||||
|
}
|
||||||
|
|
||||||
|
msg = existing.HandshakePacket[2]
|
||||||
|
f.messageMetrics.Tx(header.Handshake, header.MessageSubType(msg[1]), 1)
|
||||||
|
if !via.IsRelayed {
|
||||||
|
err := f.outside.WriteTo(msg, via.UdpAddr)
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithField("vpnAddrs", existing.vpnAddrs).WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).WithField("cached", true).
|
||||||
|
WithError(err).Error("Failed to send handshake message")
|
||||||
|
} else {
|
||||||
|
f.l.WithField("vpnAddrs", existing.vpnAddrs).WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).WithField("cached", true).
|
||||||
|
Info("Handshake message sent")
|
||||||
|
}
|
||||||
|
return
|
||||||
|
} else {
|
||||||
|
if via.relay == nil {
|
||||||
|
f.l.Error("Handshake send failed: both addr and via.relay are nil.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||||
|
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||||
|
f.l.WithField("vpnAddrs", existing.vpnAddrs).WithField("relay", via.relayHI.vpnAddrs[0]).
|
||||||
|
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).WithField("cached", true).
|
||||||
|
Info("Handshake message sent")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
case ErrExistingHostInfo:
|
||||||
|
// This means there was an existing tunnel and this handshake was older than the one we are currently based on
|
||||||
|
f.l.WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("oldHandshakeTime", existing.lastHandshakeTime).
|
||||||
|
WithField("newHandshakeTime", hostinfo.lastHandshakeTime).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
WithField("issuer", issuer).
|
||||||
|
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||||
|
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||||
|
Info("Handshake too old")
|
||||||
|
|
||||||
|
// Send a test packet to trigger an authenticated tunnel test, this should suss out any lingering tunnel issues
|
||||||
|
f.SendMessageToVpnAddr(header.Test, header.TestRequest, vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
||||||
|
return
|
||||||
|
case ErrLocalIndexCollision:
|
||||||
|
// This means we failed to insert because of collision on localIndexId. Just let the next handshake packet retry
|
||||||
|
f.l.WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
WithField("issuer", issuer).
|
||||||
|
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||||
|
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||||
|
WithField("localIndex", hostinfo.localIndexId).WithField("collision", existing.vpnAddrs).
|
||||||
|
Error("Failed to add HostInfo due to localIndex collision")
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
// Shouldn't happen, but just in case someone adds a new error type to CheckAndComplete
|
||||||
|
// And we forget to update it here
|
||||||
|
f.l.WithError(err).WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
WithField("issuer", issuer).
|
||||||
|
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||||
|
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||||
|
Error("Failed to add HostInfo to HostMap")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Do the send
|
||||||
|
f.messageMetrics.Tx(header.Handshake, header.MessageSubType(msg[1]), 1)
|
||||||
|
if !via.IsRelayed {
|
||||||
|
err = f.outside.WriteTo(msg, via.UdpAddr)
|
||||||
|
log := f.l.WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
WithField("issuer", issuer).
|
||||||
|
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||||
|
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 2, "style": "ix_psk0"})
|
||||||
|
if err != nil {
|
||||||
|
log.WithError(err).Error("Failed to send handshake")
|
||||||
|
} else {
|
||||||
|
log.Info("Handshake message sent")
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if via.relay == nil {
|
||||||
|
f.l.Error("Handshake send failed: both addr and via.relay are nil.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||||
|
// I successfully received a handshake. Just in case I marked this tunnel as 'Disestablished', ensure
|
||||||
|
// it's correctly marked as working.
|
||||||
|
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
|
||||||
|
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||||
|
f.l.WithField("vpnAddrs", vpnAddrs).WithField("relay", via.relayHI.vpnAddrs[0]).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
WithField("issuer", issuer).
|
||||||
|
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||||
|
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||||
|
Info("Handshake message sent")
|
||||||
|
}
|
||||||
|
|
||||||
|
f.connectionManager.AddTrafficWatch(hostinfo)
|
||||||
|
|
||||||
|
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
func ixHandshakeStage2(f *Interface, via ViaSender, hh *HandshakeHostInfo, packet []byte, h *header.H) bool {
|
||||||
|
if hh == nil {
|
||||||
|
// Nothing here to tear down, got a bogus stage 2 packet
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
hh.Lock()
|
||||||
|
defer hh.Unlock()
|
||||||
|
|
||||||
|
hostinfo := hh.hostinfo
|
||||||
|
if !via.IsRelayed {
|
||||||
|
// The vpnAddr we know about is the one we tried to handshake with, use it to apply the remote allow list.
|
||||||
|
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
|
||||||
|
f.l.WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).Debug("lighthouse.remote_allow_list denied incoming handshake")
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
ci := hostinfo.ConnectionState
|
||||||
|
msg, eKey, dKey, err := ci.H.ReadMessage(nil, packet[header.Len:])
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).WithField("header", h).
|
||||||
|
Error("Failed to call noise.ReadMessage")
|
||||||
|
|
||||||
|
// We don't want to tear down the connection on a bad ReadMessage because it could be an attacker trying
|
||||||
|
// to DOS us. Every other error condition after should to allow a possible good handshake to complete in the
|
||||||
|
// near future
|
||||||
|
return false
|
||||||
|
} else if dKey == nil || eKey == nil {
|
||||||
|
f.l.WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||||
|
Error("Noise did not arrive at a key")
|
||||||
|
|
||||||
|
// This should be impossible in IX but just in case, if we get here then there is no chance to recover
|
||||||
|
// the handshake state machine. Tear it down
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
hs := &NebulaHandshake{}
|
||||||
|
err = hs.Unmarshal(msg)
|
||||||
|
if err != nil || hs.Details == nil {
|
||||||
|
f.l.WithError(err).WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).Error("Failed unmarshal handshake message")
|
||||||
|
|
||||||
|
// The handshake state machine is complete, if things break now there is no chance to recover. Tear down and start again
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
rc, err := cert.Recombine(cert.Version(hs.Details.CertVersion), hs.Details.Cert, ci.H.PeerStatic(), ci.Curve())
|
||||||
|
if err != nil {
|
||||||
|
f.l.WithError(err).WithField("from", via).
|
||||||
|
WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||||
|
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||||
|
Info("Handshake did not contain a certificate")
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
remoteCert, err := f.pki.GetCAPool().VerifyCertificate(time.Now(), rc)
|
||||||
|
if err != nil {
|
||||||
|
fp, err := rc.Fingerprint()
|
||||||
|
if err != nil {
|
||||||
|
fp = "<error generating certificate fingerprint>"
|
||||||
|
}
|
||||||
|
|
||||||
|
e := f.l.WithError(err).WithField("from", via).
|
||||||
|
WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||||
|
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||||
|
WithField("certFingerprint", fp).
|
||||||
|
WithField("certVpnNetworks", rc.Networks())
|
||||||
|
|
||||||
|
if f.l.Level >= logrus.DebugLevel {
|
||||||
|
e = e.WithField("cert", rc)
|
||||||
|
}
|
||||||
|
|
||||||
|
e.Info("Invalid certificate from host")
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if !bytes.Equal(remoteCert.Certificate.PublicKey(), ci.H.PeerStatic()) {
|
||||||
|
f.l.WithField("from", via).
|
||||||
|
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||||
|
WithField("cert", remoteCert).Info("public key mismatch between certificate and handshake")
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(remoteCert.Certificate.Networks()) == 0 {
|
||||||
|
f.l.WithError(err).WithField("from", via).
|
||||||
|
WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||||
|
WithField("cert", remoteCert).
|
||||||
|
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||||
|
Info("No networks in certificate")
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
vpnNetworks := remoteCert.Certificate.Networks()
|
||||||
|
certName := remoteCert.Certificate.Name()
|
||||||
|
certVersion := remoteCert.Certificate.Version()
|
||||||
|
fingerprint := remoteCert.Fingerprint
|
||||||
|
issuer := remoteCert.Certificate.Issuer()
|
||||||
|
|
||||||
|
hostinfo.remoteIndexId = hs.Details.ResponderIndex
|
||||||
|
hostinfo.lastHandshakeTime = hs.Details.Time
|
||||||
|
|
||||||
|
// Store their cert and our symmetric keys
|
||||||
|
ci.peerCert = remoteCert
|
||||||
|
ci.dKey = NewNebulaCipherState(dKey)
|
||||||
|
ci.eKey = NewNebulaCipherState(eKey)
|
||||||
|
|
||||||
|
// Make sure the current udpAddr being used is set for responding
|
||||||
|
if !via.IsRelayed {
|
||||||
|
hostinfo.SetRemote(via.UdpAddr)
|
||||||
|
} else {
|
||||||
|
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
correctHostResponded := false
|
||||||
|
anyVpnAddrsInCommon := false
|
||||||
|
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
||||||
|
for i, network := range vpnNetworks {
|
||||||
|
vpnAddrs[i] = network.Addr()
|
||||||
|
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
||||||
|
anyVpnAddrsInCommon = true
|
||||||
|
}
|
||||||
|
if hostinfo.vpnAddrs[0] == network.Addr() {
|
||||||
|
// todo is it more correct to see if any of hostinfo.vpnAddrs are in the cert? it should have len==1, but one day it might not?
|
||||||
|
correctHostResponded = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure the right host responded
|
||||||
|
if !correctHostResponded {
|
||||||
|
f.l.WithField("intendedVpnAddrs", hostinfo.vpnAddrs).WithField("haveVpnNetworks", vpnNetworks).
|
||||||
|
WithField("from", via).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||||
|
Info("Incorrect host responded to handshake")
|
||||||
|
|
||||||
|
// Release our old handshake from pending, it should not continue
|
||||||
|
f.handshakeManager.DeleteHostInfo(hostinfo)
|
||||||
|
|
||||||
|
// Create a new hostinfo/handshake for the intended vpn ip
|
||||||
|
//TODO is hostinfo.vpnAddrs[0] always the address to use?
|
||||||
|
f.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], func(newHH *HandshakeHostInfo) {
|
||||||
|
// Block the current used address
|
||||||
|
newHH.hostinfo.remotes = hostinfo.remotes
|
||||||
|
newHH.hostinfo.remotes.BlockRemote(via)
|
||||||
|
|
||||||
|
f.l.WithField("blockedUdpAddrs", newHH.hostinfo.remotes.CopyBlockedRemotes()).
|
||||||
|
WithField("vpnNetworks", vpnNetworks).
|
||||||
|
WithField("remotes", newHH.hostinfo.remotes.CopyAddrs(f.hostMap.GetPreferredRanges())).
|
||||||
|
Info("Blocked addresses for handshakes")
|
||||||
|
|
||||||
|
// Swap the packet store to benefit the original intended recipient
|
||||||
|
newHH.packetStore = hh.packetStore
|
||||||
|
hh.packetStore = []*cachedPacket{}
|
||||||
|
|
||||||
|
// Finally, put the correct vpn addrs in the host info, tell them to close the tunnel, and return true to tear down
|
||||||
|
hostinfo.vpnAddrs = vpnAddrs
|
||||||
|
f.sendCloseTunnel(hostinfo)
|
||||||
|
})
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Mark packet 2 as seen so it doesn't show up as missed
|
||||||
|
ci.window.Update(f.l, 2)
|
||||||
|
|
||||||
|
duration := time.Since(hh.startTime).Nanoseconds()
|
||||||
|
msgRxL := f.l.WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||||
|
WithField("certName", certName).
|
||||||
|
WithField("certVersion", certVersion).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
WithField("issuer", issuer).
|
||||||
|
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||||
|
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||||
|
WithField("durationNs", duration).
|
||||||
|
WithField("sentCachedPackets", len(hh.packetStore))
|
||||||
|
if anyVpnAddrsInCommon {
|
||||||
|
msgRxL.Info("Handshake message received")
|
||||||
|
} else {
|
||||||
|
//todo warn if not lighthouse or relay?
|
||||||
|
msgRxL.Info("Handshake message received, but no vpnNetworks in common.")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build up the radix for the firewall if we have subnets in the cert
|
||||||
|
hostinfo.vpnAddrs = vpnAddrs
|
||||||
|
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
||||||
|
|
||||||
|
// Complete our handshake and update metrics, this will replace any existing tunnels for the vpnAddrs here
|
||||||
|
f.handshakeManager.Complete(hostinfo, f)
|
||||||
|
f.connectionManager.AddTrafficWatch(hostinfo)
|
||||||
|
|
||||||
|
if f.l.Level >= logrus.DebugLevel {
|
||||||
|
hostinfo.logger(f.l).Debugf("Sending %d stored packets", len(hh.packetStore))
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(hh.packetStore) > 0 {
|
||||||
|
nb := make([]byte, 12, 12)
|
||||||
|
out := make([]byte, mtu)
|
||||||
|
for _, cp := range hh.packetStore {
|
||||||
|
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
||||||
|
}
|
||||||
|
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||||
|
f.metricHandshakes.Update(duration)
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
+222
-661
File diff suppressed because it is too large
Load Diff
+1
-136
@@ -5,7 +5,6 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
@@ -28,7 +27,7 @@ func Test_NewHandshakeManagerVpnIp(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
blah := NewHandshakeManager(l, mainHM, lh, &udp.NoopConn{}, defaultHandshakeConfig)
|
blah := NewHandshakeManager(l, mainHM, lh, &udp.NoopConn{}, defaultHandshakeConfig)
|
||||||
@@ -101,137 +100,3 @@ func (mw *mockEncWriter) GetHostInfo(_ netip.Addr) *HostInfo {
|
|||||||
func (mw *mockEncWriter) GetCertState() *CertState {
|
func (mw *mockEncWriter) GetCertState() *CertState {
|
||||||
return &CertState{initiatingVersion: cert.Version2}
|
return &CertState{initiatingVersion: cert.Version2}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestValidatePeerCert(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
|
|
||||||
myNetwork := netip.MustParsePrefix("10.0.0.1/24")
|
|
||||||
myAddrTable := new(bart.Lite)
|
|
||||||
myAddrTable.Insert(netip.PrefixFrom(myNetwork.Addr(), myNetwork.Addr().BitLen()))
|
|
||||||
myNetTable := new(bart.Lite)
|
|
||||||
myNetTable.Insert(myNetwork.Masked())
|
|
||||||
|
|
||||||
newHM := func() *HandshakeManager {
|
|
||||||
hm := NewHandshakeManager(l, newHostMap(l), newTestLighthouse(), &udp.NoopConn{}, defaultHandshakeConfig)
|
|
||||||
hm.f = &Interface{
|
|
||||||
handshakeManager: hm,
|
|
||||||
pki: &PKI{},
|
|
||||||
l: l,
|
|
||||||
myVpnAddrsTable: myAddrTable,
|
|
||||||
myVpnNetworksTable: myNetTable,
|
|
||||||
lightHouse: hm.lightHouse,
|
|
||||||
}
|
|
||||||
return hm
|
|
||||||
}
|
|
||||||
|
|
||||||
cached := func(networks ...netip.Prefix) *cert.CachedCertificate {
|
|
||||||
return &cert.CachedCertificate{
|
|
||||||
Certificate: &dummyCert{name: "peer", networks: networks},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
via := ViaSender{
|
|
||||||
UdpAddr: netip.MustParseAddrPort("198.51.100.7:4242"),
|
|
||||||
IsRelayed: true, // skip the remote allow list (covered separately)
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Run("addr inside our networks sets anyVpnAddrsInCommon", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
// 10.0.0.2 falls inside our 10.0.0.0/24
|
|
||||||
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("10.0.0.2/24")))
|
|
||||||
assert.True(t, ok)
|
|
||||||
assert.True(t, common)
|
|
||||||
assert.Equal(t, []netip.Addr{netip.MustParseAddr("10.0.0.2")}, addrs)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("addr outside our networks leaves anyVpnAddrsInCommon false", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("192.168.1.5/24")))
|
|
||||||
assert.True(t, ok)
|
|
||||||
assert.False(t, common)
|
|
||||||
assert.Equal(t, []netip.Addr{netip.MustParseAddr("192.168.1.5")}, addrs)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("any matching network is enough", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
addrs, common, ok := hm.validatePeerCert(via, cached(
|
|
||||||
netip.MustParsePrefix("192.168.1.5/24"),
|
|
||||||
netip.MustParsePrefix("10.0.0.42/24"),
|
|
||||||
))
|
|
||||||
assert.True(t, ok)
|
|
||||||
assert.True(t, common)
|
|
||||||
assert.Len(t, addrs, 2)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("self-handshake is rejected", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
// 10.0.0.1 is in myVpnAddrsTable
|
|
||||||
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("10.0.0.1/24")))
|
|
||||||
assert.False(t, ok)
|
|
||||||
assert.False(t, common)
|
|
||||||
assert.Nil(t, addrs)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("cert with no networks is rejected", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
addrs, common, ok := hm.validatePeerCert(via, cached())
|
|
||||||
assert.False(t, ok)
|
|
||||||
assert.False(t, common)
|
|
||||||
assert.Nil(t, addrs)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandleIncomingDispatch(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
|
|
||||||
newHM := func() *HandshakeManager {
|
|
||||||
hm := NewHandshakeManager(l, newHostMap(l), newTestLighthouse(), &udp.NoopConn{}, defaultHandshakeConfig)
|
|
||||||
hm.f = &Interface{
|
|
||||||
handshakeManager: hm,
|
|
||||||
pki: &PKI{},
|
|
||||||
l: l,
|
|
||||||
}
|
|
||||||
return hm
|
|
||||||
}
|
|
||||||
|
|
||||||
via := ViaSender{
|
|
||||||
UdpAddr: netip.MustParseAddrPort("198.51.100.7:4242"),
|
|
||||||
IsRelayed: true, // bypass remote allow list
|
|
||||||
}
|
|
||||||
|
|
||||||
// A packet body of zero length is fine for these tests: dispatch is
|
|
||||||
// gated on header fields, and we assert that we never reach noise/cert
|
|
||||||
// processing for any of the malformed shapes here.
|
|
||||||
pkt := make([]byte, header.Len)
|
|
||||||
|
|
||||||
t.Run("unsupported subtype dropped", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
h := &header.H{Type: header.Handshake, Subtype: header.MessageSubType(99), MessageCounter: 1}
|
|
||||||
hm.HandleIncoming(via, pkt, h)
|
|
||||||
assert.Empty(t, hm.indexes, "no pending handshake should be created")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("stage-1 with non-zero RemoteIndex dropped", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
h := &header.H{
|
|
||||||
Type: header.Handshake,
|
|
||||||
Subtype: header.HandshakeIXPSK0,
|
|
||||||
RemoteIndex: 0xdeadbeef,
|
|
||||||
MessageCounter: 1,
|
|
||||||
}
|
|
||||||
hm.HandleIncoming(via, pkt, h)
|
|
||||||
assert.Empty(t, hm.indexes, "spoofed stage-1 must not create a pending machine")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("continuation with no matching pending index dropped", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
h := &header.H{
|
|
||||||
Type: header.Handshake,
|
|
||||||
Subtype: header.HandshakeIXPSK0,
|
|
||||||
RemoteIndex: 0xcafef00d,
|
|
||||||
MessageCounter: 2,
|
|
||||||
}
|
|
||||||
hm.HandleIncoming(via, pkt, h)
|
|
||||||
assert.Empty(t, hm.indexes, "orphan stage-2 must not create state")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -174,10 +174,6 @@ func (h *H) SubTypeName() string {
|
|||||||
return SubTypeName(h.Type, h.Subtype)
|
return SubTypeName(h.Type, h.Subtype)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *H) IsValidSubType() bool {
|
|
||||||
return IsValidSubType(h.Type, h.Subtype)
|
|
||||||
}
|
|
||||||
|
|
||||||
// SubTypeName will transform a nebula message sub type into a human string
|
// SubTypeName will transform a nebula message sub type into a human string
|
||||||
func SubTypeName(t MessageType, s MessageSubType) string {
|
func SubTypeName(t MessageType, s MessageSubType) string {
|
||||||
if n, ok := subTypeMap[t]; ok {
|
if n, ok := subTypeMap[t]; ok {
|
||||||
@@ -189,16 +185,6 @@ func SubTypeName(t MessageType, s MessageSubType) string {
|
|||||||
return "unknown"
|
return "unknown"
|
||||||
}
|
}
|
||||||
|
|
||||||
func IsValidSubType(t MessageType, s MessageSubType) bool {
|
|
||||||
if n, ok := subTypeMap[t]; ok {
|
|
||||||
if _, ok := (*n)[s]; ok {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewHeader turns bytes into a header
|
// NewHeader turns bytes into a header
|
||||||
func NewHeader(b []byte) (*H, error) {
|
func NewHeader(b []byte) (*H, error) {
|
||||||
h := new(H)
|
h := new(H)
|
||||||
|
|||||||
+148
-218
@@ -1,11 +1,9 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"slices"
|
"slices"
|
||||||
@@ -15,10 +13,10 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const defaultPromoteEvery = 1000 // Count of packets sent before we try moving a tunnel to a preferred underlay ip address
|
const defaultPromoteEvery = 1000 // Count of packets sent before we try moving a tunnel to a preferred underlay ip address
|
||||||
@@ -56,22 +54,13 @@ type Relay struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type HostMap struct {
|
type HostMap struct {
|
||||||
sync.RWMutex //Because we concurrently read and write to our maps
|
sync.RWMutex //Because we concurrently read and write to our maps
|
||||||
Indexes map[uint32]*HostInfo
|
Indexes map[uint32]*HostInfo
|
||||||
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
|
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
|
||||||
RemoteIndexes map[uint32]*HostInfo
|
RemoteIndexes map[uint32]*HostInfo
|
||||||
// Hosts maps a vpn address to its primary hostinfo, one entry per address we hold a tunnel
|
|
||||||
// for. moreHosts only has an entry while an address is held by 2 or more hostinfos and stores
|
|
||||||
// the full most-recent-first list; moreHosts[a][0] is always the same hostinfo as Hosts[a].
|
|
||||||
// Each address gets its own independent list, so a hostinfo owning multiple addresses can
|
|
||||||
// never corrupt another address's ordering the way the old shared next/prev chain could.
|
|
||||||
// Entries in moreHosts are only ever written by unlockedSetHostsForAddr; Hosts is written
|
|
||||||
// directly only in the single-hostinfo fast paths where moreHosts is known to have no entry,
|
|
||||||
// and unlockedDeleteHostInfo swaps either map for a fresh one when it fully drains.
|
|
||||||
Hosts map[netip.Addr]*HostInfo
|
Hosts map[netip.Addr]*HostInfo
|
||||||
moreHosts map[netip.Addr][]*HostInfo
|
|
||||||
preferredRanges atomic.Pointer[[]netip.Prefix]
|
preferredRanges atomic.Pointer[[]netip.Prefix]
|
||||||
l *slog.Logger
|
l *logrus.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
// For synchronization, treat the pointed-to Relay struct as immutable. To edit the Relay
|
// For synchronization, treat the pointed-to Relay struct as immutable. To edit the Relay
|
||||||
@@ -147,9 +136,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
|
||||||
}
|
}
|
||||||
@@ -238,7 +227,7 @@ const (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type HostInfo struct {
|
type HostInfo struct {
|
||||||
remote atomic.Pointer[netip.AddrPort]
|
remote netip.AddrPort
|
||||||
remotes *RemoteList
|
remotes *RemoteList
|
||||||
promoteCounter atomic.Uint32
|
promoteCounter atomic.Uint32
|
||||||
ConnectionState *ConnectionState
|
ConnectionState *ConnectionState
|
||||||
@@ -275,6 +264,10 @@ 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
|
||||||
|
|
||||||
@@ -287,6 +280,7 @@ 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
|
||||||
}
|
}
|
||||||
@@ -319,7 +313,7 @@ type cachedPacketMetrics struct {
|
|||||||
dropped metrics.Counter
|
dropped metrics.Counter
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewHostMapFromConfig(l *slog.Logger, c *config.C) *HostMap {
|
func NewHostMapFromConfig(l *logrus.Logger, c *config.C) *HostMap {
|
||||||
hm := newHostMap(l)
|
hm := newHostMap(l)
|
||||||
|
|
||||||
hm.reload(c, true)
|
hm.reload(c, true)
|
||||||
@@ -327,18 +321,18 @@ func NewHostMapFromConfig(l *slog.Logger, c *config.C) *HostMap {
|
|||||||
hm.reload(c, false)
|
hm.reload(c, false)
|
||||||
})
|
})
|
||||||
|
|
||||||
l.Info("Main HostMap created", "preferredRanges", hm.GetPreferredRanges())
|
l.WithField("preferredRanges", hm.GetPreferredRanges()).
|
||||||
|
Info("Main HostMap created")
|
||||||
|
|
||||||
return hm
|
return hm
|
||||||
}
|
}
|
||||||
|
|
||||||
func newHostMap(l *slog.Logger) *HostMap {
|
func newHostMap(l *logrus.Logger) *HostMap {
|
||||||
return &HostMap{
|
return &HostMap{
|
||||||
Indexes: map[uint32]*HostInfo{},
|
Indexes: map[uint32]*HostInfo{},
|
||||||
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,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -352,10 +346,7 @@ func (hm *HostMap) reload(c *config.C, initial bool) {
|
|||||||
preferredRange, err := netip.ParsePrefix(rawPreferredRange)
|
preferredRange, err := netip.ParsePrefix(rawPreferredRange)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hm.l.Warn("Failed to parse preferred ranges, ignoring",
|
hm.l.WithError(err).WithField("range", rawPreferredRanges).Warn("Failed to parse preferred ranges, ignoring")
|
||||||
"error", err,
|
|
||||||
"range", rawPreferredRanges,
|
|
||||||
)
|
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -364,10 +355,7 @@ func (hm *HostMap) reload(c *config.C, initial bool) {
|
|||||||
|
|
||||||
oldRanges := hm.preferredRanges.Swap(&preferredRanges)
|
oldRanges := hm.preferredRanges.Swap(&preferredRanges)
|
||||||
if !initial {
|
if !initial {
|
||||||
hm.l.Info("preferred_ranges changed",
|
hm.l.WithField("oldPreferredRanges", *oldRanges).WithField("newPreferredRanges", preferredRanges).Info("preferred_ranges changed")
|
||||||
"oldPreferredRanges", *oldRanges,
|
|
||||||
"newPreferredRanges", preferredRanges,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -387,55 +375,13 @@ func (hm *HostMap) EmitStats() {
|
|||||||
metrics.GetOrRegisterGauge("hostmap.main.relayIndexes", nil).Update(int64(relaysLen))
|
metrics.GetOrRegisterGauge("hostmap.main.relayIndexes", nil).Update(int64(relaysLen))
|
||||||
}
|
}
|
||||||
|
|
||||||
// unlockedSetHostsForAddr stores the per-address hostinfo list (list[0] is the primary). An empty
|
// DeleteHostInfo will fully unlink the hostinfo and return true if it was the final hostinfo for this vpn ip
|
||||||
// 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()
|
||||||
final := hm.unlockedDeleteHostInfo(hostinfo)
|
// If we have a previous or next hostinfo then we are not the last one for this vpn ip
|
||||||
|
final := (hostinfo.next == nil && hostinfo.prev == nil)
|
||||||
|
hm.unlockedDeleteHostInfo(hostinfo)
|
||||||
hm.Unlock()
|
hm.Unlock()
|
||||||
|
|
||||||
return final
|
return final
|
||||||
@@ -447,66 +393,85 @@ func (hm *HostMap) MakePrimary(hostinfo *HostInfo) {
|
|||||||
hm.unlockedMakePrimary(hostinfo)
|
hm.unlockedMakePrimary(hostinfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
// unlockedMakePrimary reports whether hostinfo is (now) the primary for each of its addresses,
|
func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) {
|
||||||
// false only when it is no longer in the hostmap at all.
|
// Get the current primary, if it exists
|
||||||
func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) bool {
|
oldHostinfo := hm.Hosts[hostinfo.vpnAddrs[0]]
|
||||||
// A hostinfo that is no longer in the hostmap must not be re-inserted here. Callers can race
|
|
||||||
// tunnel teardown, deciding to promote under the read lock and only taking the write lock
|
// Every address in the hostinfo gets elevated to primary
|
||||||
// after a delete fully unlinked the hostinfo (connection manager swapPrimary, AddRelay). Every
|
for _, vpnAddr := range hostinfo.vpnAddrs {
|
||||||
// live hostinfo is registered in Indexes by unlockedAddHostInfo, so this is a membership test.
|
//NOTE: It is possible that we leave a dangling hostinfo here but connection manager works on
|
||||||
if hm.Indexes[hostinfo.localIndexId] != hostinfo {
|
// indexes so it should be fine.
|
||||||
return false
|
hm.Hosts[vpnAddr] = hostinfo
|
||||||
}
|
}
|
||||||
|
|
||||||
// Move hostinfo to the front (primary) of each of its address lists. The lists are
|
// If we are already primary then we won't bother re-linking
|
||||||
// independent per address, so this can never leave a dangling entry the way promoting
|
if oldHostinfo == hostinfo {
|
||||||
// against a single shared chain could.
|
return
|
||||||
for _, addr := range hostinfo.vpnAddrs {
|
|
||||||
if hm.Hosts[addr] == hostinfo {
|
|
||||||
// Already primary for this address, the list is already in the right order
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
list := removeHostInfo(hm.unlockedGetHostList(addr), hostinfo)
|
|
||||||
list = append([]*HostInfo{hostinfo}, list...)
|
|
||||||
hm.unlockedSetHostsForAddr(addr, list)
|
|
||||||
}
|
}
|
||||||
return true
|
|
||||||
|
// 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
|
||||||
}
|
}
|
||||||
|
|
||||||
// unlockedDeleteHostInfo removes hostinfo from every one of its address lists and from the index
|
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
||||||
// maps. It returns true if this was the last hostinfo for all of its addresses (we no longer have
|
|
||||||
// any tunnel to the peer), which the caller uses to decide whether to clear learned lighthouse
|
|
||||||
// state and disestablish relays.
|
|
||||||
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 {
|
for _, addr := range hostinfo.vpnAddrs {
|
||||||
if list, ok := hm.moreHosts[addr]; ok {
|
h := hm.Hosts[addr]
|
||||||
list = removeHostInfo(list, hostinfo)
|
for h != nil {
|
||||||
hm.unlockedSetHostsForAddr(addr, list)
|
if h == hostinfo {
|
||||||
if len(list) > 0 {
|
hm.unlockedInnerDeleteHostInfo(h, addr)
|
||||||
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)
|
|
||||||
} else {
|
|
||||||
// We don't hold this address but another hostinfo does, we still have a tunnel to the peer
|
|
||||||
final = false
|
|
||||||
}
|
}
|
||||||
|
h = h.next
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (hm *HostMap) unlockedInnerDeleteHostInfo(hostinfo *HostInfo, addr netip.Addr) {
|
||||||
|
primary, ok := hm.Hosts[addr]
|
||||||
|
isLastHostinfo := hostinfo.next == nil && hostinfo.prev == nil
|
||||||
|
if ok && primary == hostinfo {
|
||||||
|
// The vpn addr pointer points to the same hostinfo as the local index id, we can remove it
|
||||||
|
delete(hm.Hosts, addr)
|
||||||
|
if len(hm.Hosts) == 0 {
|
||||||
|
hm.Hosts = map[netip.Addr]*HostInfo{}
|
||||||
|
}
|
||||||
|
|
||||||
|
if hostinfo.next != nil {
|
||||||
|
// 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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Go maps never shrink their buckets, replace fully drained maps so a node that churned
|
hostinfo.next = nil
|
||||||
// through a large peer count gives the memory back. Same idiom as the index maps below.
|
hostinfo.prev = nil
|
||||||
if len(hm.Hosts) == 0 {
|
|
||||||
hm.Hosts = map[netip.Addr]*HostInfo{}
|
|
||||||
}
|
|
||||||
if len(hm.moreHosts) == 0 {
|
|
||||||
hm.moreHosts = map[netip.Addr][]*HostInfo{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// The remote index uses index ids outside our control so lets make sure we are only removing
|
// The remote index uses index ids outside our control so lets make sure we are only removing
|
||||||
// the remote index pointer here if it points to the hostinfo we are deleting
|
// the remote index pointer here if it points to the hostinfo we are deleting
|
||||||
@@ -523,14 +488,13 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
|||||||
hm.Indexes = map[uint32]*HostInfo{}
|
hm.Indexes = map[uint32]*HostInfo{}
|
||||||
}
|
}
|
||||||
|
|
||||||
if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if hm.l.Level >= logrus.DebugLevel {
|
||||||
hm.l.Debug("Hostmap hostInfo deleted",
|
hm.l.WithField("hostMap", m{"mapTotalSize": len(hm.Hosts),
|
||||||
"hostMap", m{"mapTotalSize": len(hm.Hosts),
|
"vpnAddrs": hostinfo.vpnAddrs, "indexNumber": hostinfo.localIndexId, "remoteIndexNumber": hostinfo.remoteIndexId}).
|
||||||
"vpnAddrs": hostinfo.vpnAddrs, "indexNumber": hostinfo.localIndexId, "remoteIndexNumber": hostinfo.remoteIndexId},
|
Debug("Hostmap hostInfo deleted")
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if final {
|
if isLastHostinfo {
|
||||||
// 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)
|
||||||
@@ -539,8 +503,6 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
|||||||
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 {
|
||||||
@@ -584,30 +546,19 @@ 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 _, targetIp := range targetIps {
|
for h != nil {
|
||||||
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
for _, targetIp := range targetIps {
|
||||||
if ok && r.State == Established {
|
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
||||||
return h, r, nil
|
if ok && r.State == Established {
|
||||||
}
|
return h, r, nil
|
||||||
}
|
|
||||||
|
|
||||||
if list, ok := hm.moreHosts[relayHostIp]; ok {
|
|
||||||
// list[0] is the primary we already checked
|
|
||||||
for _, h := range list[1:] {
|
|
||||||
for _, targetIp := range targetIps {
|
|
||||||
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
|
||||||
if ok && r.State == Established {
|
|
||||||
return h, r, nil
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
h = h.next
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil, nil, errors.New("unable to find host with relay")
|
return nil, nil, errors.New("unable to find host with relay")
|
||||||
@@ -615,14 +566,20 @@ 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() {
|
||||||
for _, h := range hm.unlockedGetHostList(relayHostIp) {
|
if h, ok := hm.Hosts[relayHostIp]; ok {
|
||||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
for h != nil {
|
||||||
|
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 {
|
||||||
for _, h := range hm.unlockedGetHostList(rs.PeerAddr) {
|
if h, ok := hm.Hosts[rs.PeerAddr]; ok {
|
||||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
for h != nil {
|
||||||
|
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
||||||
|
h = h.next
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -658,41 +615,30 @@ 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 hm.l.Level >= logrus.DebugLevel {
|
||||||
if f.connectionManager != nil { // f.connectionManager is only nil in some unit tests
|
hm.l.WithField("hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts),
|
||||||
f.connectionManager.trafficTimer.Add(hostinfo.localIndexId, f.connectionManager.checkInterval)
|
"hostinfo": m{"existing": true, "localIndexId": hostinfo.localIndexId, "vpnAddrs": hostinfo.vpnAddrs}}).
|
||||||
}
|
Debug("Hostmap vpnIp added")
|
||||||
|
|
||||||
if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hm.l.Debug("Hostmap vpnIp added",
|
|
||||||
"hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts),
|
|
||||||
"hostinfo": m{"existing": true, "localIndexId": hostinfo.localIndexId, "vpnAddrs": hostinfo.vpnAddrs}},
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) {
|
func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) {
|
||||||
existing, ok := hm.Hosts[vpnAddr]
|
existing := hm.Hosts[vpnAddr]
|
||||||
if !ok {
|
hm.Hosts[vpnAddr] = hostinfo
|
||||||
// Common case, the first hostinfo for this address. moreHosts stays empty.
|
|
||||||
hm.Hosts[vpnAddr] = hostinfo
|
if existing != nil && existing != hostinfo {
|
||||||
return
|
hostinfo.next = existing
|
||||||
|
existing.prev = hostinfo
|
||||||
}
|
}
|
||||||
|
|
||||||
// The new hostinfo becomes the primary for this address. Remove any stale copy of it first so
|
i := 1
|
||||||
// we never hold a duplicate, then prepend.
|
check := hostinfo
|
||||||
list, ok := hm.moreHosts[vpnAddr]
|
for check != nil {
|
||||||
if !ok {
|
if i > MaxHostInfosPerVpnIp {
|
||||||
list = []*HostInfo{existing}
|
hm.unlockedDeleteHostInfo(check)
|
||||||
}
|
}
|
||||||
list = removeHostInfo(list, hostinfo)
|
check = check.next
|
||||||
list = append([]*HostInfo{hostinfo}, list...)
|
i++
|
||||||
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])
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -724,7 +670,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.GetRemote()
|
remote := i.remote
|
||||||
|
|
||||||
// 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() {
|
||||||
@@ -766,18 +712,11 @@ 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.GetRemote() != remote {
|
if i.remote != remote {
|
||||||
i.remote.Store(&remote)
|
i.remote = remote
|
||||||
i.remotes.LearnRemote(i.vpnAddrs[0], remote)
|
i.remotes.LearnRemote(i.vpnAddrs[0], remote)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -789,7 +728,7 @@ func (i *HostInfo) SetRemoteIfPreferred(hm *HostMap, via ViaSender) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
currentRemote := i.GetRemote()
|
currentRemote := i.remote
|
||||||
if !currentRemote.IsValid() {
|
if !currentRemote.IsValid() {
|
||||||
i.SetRemote(via.UdpAddr)
|
i.SetRemote(via.UdpAddr)
|
||||||
return true
|
return true
|
||||||
@@ -845,21 +784,18 @@ func (i *HostInfo) buildNetworks(myVpnNetworksTable *bart.Lite, c cert.Certifica
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// logger returns a derived slog.Logger with per-hostinfo fields pre-bound.
|
func (i *HostInfo) logger(l *logrus.Logger) *logrus.Entry {
|
||||||
func (i *HostInfo) logger(l *slog.Logger) *slog.Logger {
|
|
||||||
if i == nil {
|
if i == nil {
|
||||||
return l
|
return logrus.NewEntry(l)
|
||||||
}
|
}
|
||||||
|
|
||||||
li := l.With(
|
li := l.WithField("vpnAddrs", i.vpnAddrs).
|
||||||
"vpnAddrs", i.vpnAddrs,
|
WithField("localIndex", i.localIndexId).
|
||||||
"localIndex", i.localIndexId,
|
WithField("remoteIndex", i.remoteIndexId)
|
||||||
"remoteIndex", i.remoteIndexId,
|
|
||||||
)
|
|
||||||
|
|
||||||
if connState := i.ConnectionState; connState != nil {
|
if connState := i.ConnectionState; connState != nil {
|
||||||
if peerCert := connState.peerCert; peerCert != nil {
|
if peerCert := connState.peerCert; peerCert != nil {
|
||||||
li = li.With("certName", peerCert.Certificate.Name())
|
li = li.WithField("certName", peerCert.Certificate.Name())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -868,17 +804,14 @@ func (i *HostInfo) logger(l *slog.Logger) *slog.Logger {
|
|||||||
|
|
||||||
// Utility functions
|
// Utility functions
|
||||||
|
|
||||||
func localAddrs(l *slog.Logger, allowList *LocalAllowList) []netip.Addr {
|
func localAddrs(l *logrus.Logger, allowList *LocalAllowList) []netip.Addr {
|
||||||
//FIXME: This function is pretty garbage
|
//FIXME: This function is pretty garbage
|
||||||
var finalAddrs []netip.Addr
|
var finalAddrs []netip.Addr
|
||||||
ifaces, _ := net.Interfaces()
|
ifaces, _ := net.Interfaces()
|
||||||
for _, i := range ifaces {
|
for _, i := range ifaces {
|
||||||
allow := allowList.AllowName(i.Name)
|
allow := allowList.AllowName(i.Name)
|
||||||
if l.Enabled(context.Background(), logging.LevelTrace) {
|
if l.Level >= logrus.TraceLevel {
|
||||||
l.Log(context.Background(), logging.LevelTrace, "localAllowList.AllowName",
|
l.WithField("interfaceName", i.Name).WithField("allow", allow).Trace("localAllowList.AllowName")
|
||||||
"interfaceName", i.Name,
|
|
||||||
"allow", allow,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if !allow {
|
if !allow {
|
||||||
@@ -896,8 +829,8 @@ func localAddrs(l *slog.Logger, allowList *LocalAllowList) []netip.Addr {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !addr.IsValid() {
|
if !addr.IsValid() {
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
if l.Level >= logrus.DebugLevel {
|
||||||
l.Debug("addr was invalid", "localAddr", rawAddr)
|
l.WithField("localAddr", rawAddr).Debug("addr was invalid")
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -905,11 +838,8 @@ func localAddrs(l *slog.Logger, allowList *LocalAllowList) []netip.Addr {
|
|||||||
|
|
||||||
if addr.IsLoopback() == false && addr.IsLinkLocalUnicast() == false {
|
if addr.IsLoopback() == false && addr.IsLinkLocalUnicast() == false {
|
||||||
isAllowed := allowList.Allow(addr)
|
isAllowed := allowList.Allow(addr)
|
||||||
if l.Enabled(context.Background(), logging.LevelTrace) {
|
if l.Level >= logrus.TraceLevel {
|
||||||
l.Log(context.Background(), logging.LevelTrace, "localAllowList.Allow",
|
l.WithField("localAddr", addr).WithField("allowed", isAllowed).Trace("localAllowList.Allow")
|
||||||
"localAddr", addr,
|
|
||||||
"allowed", isAllowed,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
if !isAllowed {
|
if !isAllowed {
|
||||||
continue
|
continue
|
||||||
|
|||||||
+139
-296
@@ -2,7 +2,6 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"slices"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
@@ -11,84 +10,78 @@ 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{a}, localIndexId: 1}
|
h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
|
||||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
|
||||||
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3}
|
h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
|
||||||
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4}
|
h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, 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)
|
||||||
|
|
||||||
// Most-recently-added is primary: h1, h2, h3, h4
|
// Make sure we go h1 -> h2 -> h3 -> h4
|
||||||
assert.Equal(t, []uint32{1, 2, 3, 4}, chainIds(t, hm, a))
|
prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||||
assert.Equal(t, h1, hm.QueryVpnAddr(a))
|
assert.Equal(t, h1.localIndexId, prim.localIndexId)
|
||||||
|
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 the middle to primary: h3, h1, h2, h4
|
// Swap h3/middle to primary
|
||||||
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))
|
|
||||||
|
|
||||||
// Swap the tail to primary: h4, h3, h1, h2
|
// Make sure we go h3 -> h1 -> h2 -> h4
|
||||||
hm.MakePrimary(h4)
|
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||||
assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a))
|
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)
|
||||||
|
|
||||||
// Swapping the current primary again is a no-op
|
// 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
|
||||||
|
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)
|
||||||
|
|
||||||
|
// 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) {
|
||||||
@@ -96,14 +89,13 @@ 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{a}, localIndexId: 1}
|
h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
|
||||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
|
||||||
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3}
|
h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
|
||||||
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4}
|
h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4}
|
||||||
h5 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 5}
|
h5 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 5}
|
||||||
h6 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 6}
|
h6 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 6}
|
||||||
|
|
||||||
hm.unlockedAddHostInfo(h6, f)
|
hm.unlockedAddHostInfo(h6, f)
|
||||||
hm.unlockedAddHostInfo(h5, f)
|
hm.unlockedAddHostInfo(h5, f)
|
||||||
@@ -112,248 +104,99 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
|
|||||||
hm.unlockedAddHostInfo(h2, f)
|
hm.unlockedAddHostInfo(h2, f)
|
||||||
hm.unlockedAddHostInfo(h1, f)
|
hm.unlockedAddHostInfo(h1, f)
|
||||||
|
|
||||||
// h6 is evicted by the MaxHostInfosPerVpnIp cap; the rest are newest-first.
|
// h6 should be deleted
|
||||||
assert.Nil(t, hm.QueryIndex(h6.localIndexId))
|
assert.Nil(t, h6.next)
|
||||||
assert.Equal(t, []uint32{1, 2, 3, 4, 5}, chainIds(t, hm, a))
|
assert.Nil(t, h6.prev)
|
||||||
|
h := hm.QueryIndex(h6.localIndexId)
|
||||||
|
assert.Nil(t, h)
|
||||||
|
|
||||||
// Delete primary; not final since siblings remain.
|
// Make sure we go h1 -> h2 -> h3 -> h4 -> h5
|
||||||
assert.False(t, hm.DeleteHostInfo(h1))
|
prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||||
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a))
|
assert.Equal(t, h1.localIndexId, prim.localIndexId)
|
||||||
|
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)
|
||||||
|
|
||||||
// Deleting the same hostinfo again must not report final while siblings remain and must not
|
// Delete primary
|
||||||
// disturb the list. The old chain code got this wrong: the first delete nil'd next/prev, so a
|
hm.DeleteHostInfo(h1)
|
||||||
// second delete looked final and wiped lighthouse state out from under the live sibling.
|
assert.Nil(t, h1.prev)
|
||||||
assert.False(t, hm.DeleteHostInfo(h1))
|
assert.Nil(t, h1.next)
|
||||||
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a))
|
|
||||||
|
|
||||||
// Delete a middle node.
|
// Make sure we go h2 -> h3 -> h4 -> h5
|
||||||
assert.False(t, hm.DeleteHostInfo(h3))
|
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||||
assert.Equal(t, []uint32{2, 4, 5}, chainIds(t, hm, a))
|
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
||||||
|
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 the tail.
|
// Delete in the middle
|
||||||
assert.False(t, hm.DeleteHostInfo(h5))
|
hm.DeleteHostInfo(h3)
|
||||||
assert.Equal(t, []uint32{2, 4}, chainIds(t, hm, a))
|
assert.Nil(t, h3.prev)
|
||||||
|
assert.Nil(t, h3.next)
|
||||||
|
|
||||||
// Delete the head; h4 remains and becomes primary.
|
// Make sure we go h2 -> h4 -> h5
|
||||||
assert.False(t, hm.DeleteHostInfo(h2))
|
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||||
assert.Equal(t, []uint32{4}, chainIds(t, hm, a))
|
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
||||||
assert.Equal(t, h4, hm.QueryVpnAddr(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.Equal(t, h5.localIndexId, h4.next.localIndexId)
|
||||||
|
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
|
||||||
|
assert.Nil(t, h5.next)
|
||||||
|
|
||||||
// Delete the only remaining item; final is true and the address is gone.
|
// Delete the tail
|
||||||
assert.True(t, hm.DeleteHostInfo(h4))
|
hm.DeleteHostInfo(h5)
|
||||||
assert.Empty(t, chainIds(t, hm, a))
|
assert.Nil(t, h5.prev)
|
||||||
assert.Nil(t, hm.QueryVpnAddr(a))
|
assert.Nil(t, h5.next)
|
||||||
|
|
||||||
// Deleting an already-gone hostinfo is still final; nothing holds the address anymore.
|
// Make sure we go h2 -> h4
|
||||||
assert.True(t, hm.DeleteHostInfo(h4))
|
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||||
assert.Empty(t, chainIds(t, hm, a))
|
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
||||||
}
|
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)
|
||||||
|
|
||||||
// TestHostMap_MakePrimary_DeletedHostInfo covers promoting a hostinfo that lost a race with
|
// Delete the head
|
||||||
// tunnel teardown: swapPrimary and AddRelay decide to promote while holding a stale pointer and
|
hm.DeleteHostInfo(h2)
|
||||||
// only take the write lock after a delete fully unlinked the hostinfo. MakePrimary must be a
|
assert.Nil(t, h2.prev)
|
||||||
// no-op, not a resurrection that installs an unmanaged primary.
|
assert.Nil(t, h2.next)
|
||||||
func TestHostMap_MakePrimary_DeletedHostInfo(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
hm := newHostMap(l)
|
|
||||||
f := &Interface{}
|
|
||||||
a := netip.MustParseAddr("0.0.0.1")
|
|
||||||
|
|
||||||
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
// Make sure we only have h4
|
||||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||||
hm.unlockedAddHostInfo(h1, f)
|
assert.Equal(t, h4.localIndexId, prim.localIndexId)
|
||||||
hm.unlockedAddHostInfo(h2, f)
|
assert.Nil(t, prim.prev)
|
||||||
|
assert.Nil(t, prim.next)
|
||||||
|
assert.Nil(t, h4.next)
|
||||||
|
|
||||||
// h1 is fully deleted while another goroutine still holds a pointer to it.
|
// Delete the only item
|
||||||
assert.False(t, hm.DeleteHostInfo(h1))
|
hm.DeleteHostInfo(h4)
|
||||||
assert.Equal(t, []uint32{2}, chainIds(t, hm, a))
|
assert.Nil(t, h4.prev)
|
||||||
|
assert.Nil(t, h4.next)
|
||||||
|
|
||||||
// The stale promote must not bring it back.
|
// Make sure we have nil
|
||||||
hm.MakePrimary(h1)
|
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||||
assert.Equal(t, []uint32{2}, chainIds(t, hm, a))
|
assert.Nil(t, prim)
|
||||||
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) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
c := config.NewC(test.NewLogger())
|
c := config.NewC(l)
|
||||||
|
|
||||||
hm := NewHostMapFromConfig(l, c)
|
hm := NewHostMapFromConfig(l, c)
|
||||||
|
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user