mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 16:37:03 +02:00
Compare commits
17 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| dfe94c6269 | |||
| 1c601d776a | |||
| 17d8ebff93 | |||
| 612d3ef931 | |||
| 8282a629e5 | |||
| c62f27d4b4 | |||
| f5db77f214 | |||
| b9a7d1edf3 | |||
| d1ea33659a | |||
| 8fdd98f639 | |||
| 45bc0fc055 | |||
| 24af30bd78 | |||
| 1d84b81032 | |||
| b155f4b7e1 | |||
| 194d58cd46 | |||
| a476b1fa07 | |||
| 8b02b8128e |
@@ -1,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
|
|
||||||
|
|||||||
@@ -18,11 +18,11 @@ 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: build
|
- name: build
|
||||||
@@ -36,14 +36,6 @@ jobs:
|
|||||||
working-directory: ./.github/workflows/smoke
|
working-directory: ./.github/workflows/smoke
|
||||||
run: ./smoke.sh
|
run: ./smoke.sh
|
||||||
|
|
||||||
- name: setup docker image ipv6
|
|
||||||
working-directory: ./.github/workflows/smoke
|
|
||||||
run: SMOKE_OVERLAY_IPV6=1 ./build.sh
|
|
||||||
|
|
||||||
- name: run smoke ipv6
|
|
||||||
working-directory: ./.github/workflows/smoke
|
|
||||||
run: SMOKE_OVERLAY_IPV6=1 ./smoke.sh
|
|
||||||
|
|
||||||
- name: setup relay docker image
|
- name: setup relay docker image
|
||||||
working-directory: ./.github/workflows/smoke
|
working-directory: ./.github/workflows/smoke
|
||||||
run: ./build-relay.sh
|
run: ./build-relay.sh
|
||||||
|
|||||||
@@ -5,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 -e '/usr/bin/echo helloagainfromhost4' -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
|
||||||
@@ -172,7 +159,7 @@ set -x
|
|||||||
# cannot initiate UDP to host2. Once host2 initiates a flow to host4:4000,
|
# cannot initiate UDP to host2. Once host2 initiates a flow to host4:4000,
|
||||||
# conntrack must let host4's listener reply on that flow. If it doesn't,
|
# conntrack must let host4's listener reply on that flow. If it doesn't,
|
||||||
# the echo back from host4 never reaches host2.
|
# the echo back from host4 never reaches host2.
|
||||||
docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv $HOST4_NIP 4000" | grep -q helloagainfromhost4
|
docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv 192.168.100.4 4000" | grep -q helloagainfromhost4
|
||||||
|
|
||||||
docker exec host4 sh -c 'kill 1'
|
docker exec host4 sh -c 'kill 1'
|
||||||
docker exec host3 sh -c 'kill 1'
|
docker exec host3 sh -c 'kill 1'
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
# -*- mode: ruby -*-
|
# -*- mode: ruby -*-
|
||||||
# vi: set ft=ruby :
|
# vi: set ft=ruby :
|
||||||
Vagrant.configure("2") do |config|
|
Vagrant.configure("2") do |config|
|
||||||
config.vm.box = "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
-97
@@ -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,109 +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 bin-boringcrypto
|
|
||||||
test-cmd: make test-boringcrypto
|
test-linux-boringcrypto:
|
||||||
e2e-cmd: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0"
|
name: Build and test on linux with boringcrypto
|
||||||
- name: linux-pkcs11
|
runs-on: ubuntu-latest
|
||||||
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
|
||||||
|
|||||||
@@ -7,88 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
## [Unreleased]
|
## [Unreleased]
|
||||||
|
|
||||||
## [1.11.0] - 2026-07-23
|
|
||||||
|
|
||||||
See the [v1.11.0](https://github.com/slackhq/nebula/milestone/25?closed=1) milestone for a complete list of changes.
|
|
||||||
|
|
||||||
### Breaking
|
|
||||||
|
|
||||||
- Logging has switched from logrus to Go's structured `slog`. Log output changes: levels are upper case
|
|
||||||
(`level=INFO`), trace prints as `level=DEBUG-4`, timestamps are always RFC3339Nano and `logging.timestamp_format`
|
|
||||||
is ignored, and some messages were reworded. Review any log parsing before upgrading. This is also an API break
|
|
||||||
for embedders, as constructors now take a `*slog.Logger`. (#1672, #1734, #1621)
|
|
||||||
- `firewall.inbound_action` and `firewall.outbound_action` (used to set reject vs. drop policy) were each being
|
|
||||||
applied to the opposite direction, that is now corrected. This only affects how blocked packets are answered, not
|
|
||||||
which packets the firewall allows or denies. If you set either of these you are getting the behavior of the other
|
|
||||||
one today and likely want to swap them before upgrading. (#1798)
|
|
||||||
- On Windows, Nebula now installs WFP PERMIT filters for the nebula adapter and the listener port by default. WFP
|
|
||||||
sits below Windows Defender Firewall, so any WDF inbound rules you rely on for either will no longer apply. Set
|
|
||||||
`tun.windows_bypass_wdf` and `listen.windows_bypass_wdf` to false to leave WDF in charge. (#1710)
|
|
||||||
- On Windows, the nebula device is now set to the `private` network category instead of whatever Windows decided,
|
|
||||||
which is usually `Public`. This makes the host firewall less restrictive on the overlay. Set
|
|
||||||
`tun.network_category` to `unset` to keep the old behavior. (#1710)
|
|
||||||
- Reject packets for non-TCP now use ICMP code 13, communication administratively prohibited, instead of code 3,
|
|
||||||
port unreachable. Anything keying off the old code needs updating. (#1766, #1768)
|
|
||||||
- The SSH debug server's profiling commands are now confined to `sshd.sandbox_dir`, which defaults to
|
|
||||||
`$TMP/nebula-debug`. Relative paths resolve inside it and absolute paths outside it are rejected, so anything
|
|
||||||
scripting `start-cpu-profile`, `save-heap-profile`, or `save-mutex-profile` with a path elsewhere needs the
|
|
||||||
directory set. The directory is not created for you. (#1622)
|
|
||||||
|
|
||||||
### Added
|
|
||||||
|
|
||||||
- Sign the Windows release binaries. (#1718)
|
|
||||||
- Generate IPv6 reject packets, matching the existing IPv4 behavior. (#1766, #1767, #1768)
|
|
||||||
- Accept `-` in `nebula-cert` to read from stdin or write to stdout. (#1714)
|
|
||||||
- Search for both `config.yml` and `config.yaml` in service and command line modes. (#1717)
|
|
||||||
- Add version labels to the Docker/OCI images. (#1772)
|
|
||||||
- Rebind the listener and re-query lighthouses on macOS when the underlay network changes, so devices moving
|
|
||||||
between wifi and wired or between networks recover without waiting for dead tunnel detection. Controlled by
|
|
||||||
`listen.rebind_on_network_change` (default `true`, not reloadable). (#1816)
|
|
||||||
|
|
||||||
### Changed
|
|
||||||
|
|
||||||
- Reload the firewall when the unsafe networks in the certificate change. (#1719)
|
|
||||||
- Reconfigure, start, and stop the stats listener on a config reload instead of requiring a restart. (#1670)
|
|
||||||
- Update a static host's addresses when they change on reload. (#1713)
|
|
||||||
- Don't require a port on ICMP firewall rules. (#1609)
|
|
||||||
- Connection track ICMP traffic. (#1602)
|
|
||||||
- Return `NODATA` instead of `NXDOMAIN` from the DNS server for a name that exists but has no record of the
|
|
||||||
requested type, so clients that query `AAAA` first (busybox/Alpine) fall through to `A`. (#1668)
|
|
||||||
- Record the local host's details in the DNS server. (#1716)
|
|
||||||
- Install Windows unsafe routes as link routes. (#1709)
|
|
||||||
- Reduce relay handshake log spam, and only log a handshake send error at error level when the remote list
|
|
||||||
changes. (#1733, #1765, #1810)
|
|
||||||
- Start, stop, and reload subsystems (DNS, stats, conntrack, ssh, punchy) cleanly without leaking goroutines. (#1640, #1654, #1661, #1667, #1669, #1708, #1806, #1815)
|
|
||||||
- `Control` is now safe to stop and wait on from any lifecycle state, and a new `Control.Wait` blocks until nebula
|
|
||||||
has fully stopped and returns the first fatal reader error. Failed starts release the udp sockets and tun fd
|
|
||||||
instead of leaking them. (#1794)
|
|
||||||
- Trigger an immediate lighthouse update when reconnecting to or adding a lighthouse instead of waiting for the next update tick. (#1645)
|
|
||||||
- Bring the Darwin and OpenBSD tun implementations in line with the other BSDs. (#1703)
|
|
||||||
- Update to build against go v1.26. (#1818)
|
|
||||||
- Various dependency updates. (#1586, #1587, #1604, #1617, #1618, #1627, #1628, #1629, #1652, #1664, #1665, #1697, #1721, #1732, #1742, #1743, #1750, #1763, #1771, #1782, #1800, #1807)
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Fix a data race on a host's remote address that could send packets to the wrong address during a roam. (#1773)
|
|
||||||
- Fix tunnels that could permanently escape connection manager monitoring. (#1752)
|
|
||||||
- Fix a crash when reloading the SSH server's trusted keys. (#1787)
|
|
||||||
- Fix hostmap corruption when a host has multiple overlay addresses. Each address now gets its own list instead of
|
|
||||||
a single shared chain, which also fixes two latent bugs on the add and makePrimary paths. (#1788, #1790)
|
|
||||||
- Apply `remote_allow_list` IPv4 rules to 4-in-6 mapped addresses. (#1786)
|
|
||||||
- Don't panic in the DNS server on a short or empty query name. (#1635)
|
|
||||||
- Advance the replay window on relayed packets so a relay drops replayed frames instead of re-forwarding them. (#1751)
|
|
||||||
- Fix a race in relay state handling. (#1753)
|
|
||||||
- Lock replay window updates so concurrent readers can't corrupt it. (#1802)
|
|
||||||
- Reject malformed handshakes more reliably, including invalid ed25519 key lengths. (#1601, #1756)
|
|
||||||
- Properly handle `closetunnel` packets. (#1638)
|
|
||||||
- Fix an IPv6 extension-header length overflow that could make the firewall parse the wrong protocol and ports. (#1789)
|
|
||||||
- Fix relay re-establishment when a handshake arrives over a relay entry that a one-sided teardown left
|
|
||||||
`Disestablished`, which silently dropped every send until dead tunnel detection forced a re-handshake. (#1805)
|
|
||||||
- Don't build new relay state on a tunnel that was just discarded. (#1796)
|
|
||||||
- Don't delete the wrong pending hostinfo in the handshake manager. (#1811)
|
|
||||||
- Don't call the packet reader after a UDP error on Darwin. (#1755)
|
|
||||||
- Open the FreeBSD tun device non blocking. (#1666)
|
|
||||||
|
|
||||||
## [1.10.3] - 2026-02-06
|
## [1.10.3] - 2026-02-06
|
||||||
|
|
||||||
### Security
|
### Security
|
||||||
|
|||||||
@@ -60,18 +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
|
|
||||||
|
|
||||||
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
|
||||||
|
|
||||||
@@ -94,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)
|
||||||
@@ -268,9 +227,6 @@ smoke-relay-docker: bin-docker
|
|||||||
cd .github/workflows/smoke/ && ./build-relay.sh
|
cd .github/workflows/smoke/ && ./build-relay.sh
|
||||||
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
||||||
|
|
||||||
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
|
|
||||||
smoke-docker-ipv6: smoke-docker
|
|
||||||
|
|
||||||
smoke-docker-race: BUILD_ARGS = -race
|
smoke-docker-race: BUILD_ARGS = -race
|
||||||
smoke-docker-race: CGO_ENABLED = 1
|
smoke-docker-race: CGO_ENABLED = 1
|
||||||
smoke-docker-race: smoke-docker
|
smoke-docker-race: smoke-docker
|
||||||
@@ -280,5 +236,5 @@ smoke-vagrant/%: bin-docker build/%/nebula
|
|||||||
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 build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html 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
|
||||||
|
|||||||
@@ -2,42 +2,24 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"math"
|
|
||||||
mathbits "math/bits"
|
|
||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
)
|
)
|
||||||
|
|
||||||
const bitsPerWord = 64
|
|
||||||
|
|
||||||
// Bits is a sliding-window anti-replay tracker. The window is stored as a
|
|
||||||
// circular bitmap packed into uint64 words (8x denser than a []bool), so a
|
|
||||||
// length-N window costs N/8 bytes. length must be a power of two.
|
|
||||||
type Bits struct {
|
type Bits struct {
|
||||||
length uint64
|
length uint64
|
||||||
lengthMask uint64
|
|
||||||
current uint64
|
current uint64
|
||||||
bits []uint64
|
bits []bool
|
||||||
lostCounter metrics.Counter
|
lostCounter metrics.Counter
|
||||||
dupeCounter metrics.Counter
|
dupeCounter metrics.Counter
|
||||||
outOfWindowCounter metrics.Counter
|
outOfWindowCounter metrics.Counter
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewBits(length uint64) *Bits {
|
func NewBits(bits uint64) *Bits {
|
||||||
if length == 0 || length&(length-1) != 0 {
|
|
||||||
panic(fmt.Sprintf("Bits length must be a power of two, got %d", length))
|
|
||||||
}
|
|
||||||
|
|
||||||
nWords := length / bitsPerWord
|
|
||||||
if nWords == 0 {
|
|
||||||
nWords = 1
|
|
||||||
}
|
|
||||||
b := &Bits{
|
b := &Bits{
|
||||||
length: length,
|
length: bits,
|
||||||
lengthMask: length - 1,
|
bits: make([]bool, bits, bits),
|
||||||
bits: make([]uint64, nWords),
|
|
||||||
current: 0,
|
current: 0,
|
||||||
lostCounter: metrics.GetOrRegisterCounter("network.packets.lost", nil),
|
lostCounter: metrics.GetOrRegisterCounter("network.packets.lost", nil),
|
||||||
dupeCounter: metrics.GetOrRegisterCounter("network.packets.duplicate", nil),
|
dupeCounter: metrics.GetOrRegisterCounter("network.packets.duplicate", nil),
|
||||||
@@ -45,194 +27,71 @@ func NewBits(length uint64) *Bits {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// There is no counter value 0, mark it to avoid counting a lost packet later.
|
// There is no counter value 0, mark it to avoid counting a lost packet later.
|
||||||
b.bits[0] = 1
|
b.bits[0] = true
|
||||||
|
b.current = 0
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *Bits) get(i uint64) bool {
|
|
||||||
pos := i & b.lengthMask
|
|
||||||
//bit-shifting by 6 because i is a bit index, not a u64 index, and we need to find the u64 without bit in it
|
|
||||||
return b.bits[pos>>6]&(uint64(1)<<(pos&63)) != 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (b *Bits) set(i uint64) {
|
|
||||||
pos := i & b.lengthMask
|
|
||||||
b.bits[pos>>6] |= uint64(1) << (pos & 63)
|
|
||||||
}
|
|
||||||
|
|
||||||
// clearRange clears `count` bits starting at circular position `startPos`
|
|
||||||
// (already masked to [0, length)) and returns how many of them were set
|
|
||||||
// before the clear. count must be in [1, length].
|
|
||||||
func (b *Bits) clearRange(startPos, count uint64) uint64 {
|
|
||||||
wasSet := uint64(0)
|
|
||||||
if count >= b.length {
|
|
||||||
for _, w := range b.bits {
|
|
||||||
wasSet += uint64(mathbits.OnesCount64(w))
|
|
||||||
}
|
|
||||||
clear(b.bits)
|
|
||||||
return wasSet
|
|
||||||
}
|
|
||||||
|
|
||||||
pos := startPos
|
|
||||||
remaining := count
|
|
||||||
|
|
||||||
// handle the potential partial word before pos becomes u64 aligned
|
|
||||||
word := pos >> 6
|
|
||||||
bit := pos & 63
|
|
||||||
take := uint64(64) - bit
|
|
||||||
if take > remaining {
|
|
||||||
take = remaining
|
|
||||||
}
|
|
||||||
if take > b.length-pos {
|
|
||||||
take = b.length - pos
|
|
||||||
}
|
|
||||||
var mask uint64
|
|
||||||
if take == 64 {
|
|
||||||
mask = math.MaxUint64
|
|
||||||
} else {
|
|
||||||
mask = ((uint64(1) << take) - 1) << bit
|
|
||||||
}
|
|
||||||
wasSet += uint64(mathbits.OnesCount64(b.bits[word] & mask))
|
|
||||||
b.bits[word] &^= mask
|
|
||||||
remaining -= take
|
|
||||||
pos = (pos + take) & b.lengthMask
|
|
||||||
|
|
||||||
// Clear whole words, keeping track of the number of set bits
|
|
||||||
for remaining >= 64 {
|
|
||||||
word = pos >> 6
|
|
||||||
wasSet += uint64(mathbits.OnesCount64(b.bits[word]))
|
|
||||||
b.bits[word] = 0
|
|
||||||
remaining -= 64
|
|
||||||
pos = (pos + 64) & b.lengthMask
|
|
||||||
}
|
|
||||||
|
|
||||||
// Clear the remaining partial word
|
|
||||||
if remaining > 0 {
|
|
||||||
word = pos >> 6
|
|
||||||
mask = (uint64(1) << remaining) - 1
|
|
||||||
wasSet += uint64(mathbits.OnesCount64(b.bits[word] & mask))
|
|
||||||
b.bits[word] &^= mask
|
|
||||||
}
|
|
||||||
|
|
||||||
return wasSet
|
|
||||||
}
|
|
||||||
|
|
||||||
func (b *Bits) strictlyWithinWindow(i uint64) bool {
|
|
||||||
// Handle the case where the window hasn't slid yet. This avoids u64 underflow.
|
|
||||||
inWarmup := b.current < b.length
|
|
||||||
if i < b.length && inWarmup {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// Next, if the packet is in-window, see if we've seen it before
|
|
||||||
if i > b.current-b.length {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
return false //not within window!
|
|
||||||
}
|
|
||||||
|
|
||||||
// Check returns true if i is within (or way out in front of) the window, and not a replay
|
|
||||||
func (b *Bits) Check(l *slog.Logger, i uint64) bool {
|
func (b *Bits) Check(l *slog.Logger, i uint64) bool {
|
||||||
// If i is the next number, return true.
|
// If i is the next number, return true.
|
||||||
if i > b.current {
|
if i > b.current {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
if b.strictlyWithinWindow(i) {
|
// If i is within the window, check if it's been set already.
|
||||||
return !b.get(i)
|
if i > b.current-b.length || i < b.length && b.current < b.length {
|
||||||
|
return !b.bits[i%b.length]
|
||||||
}
|
}
|
||||||
|
|
||||||
// Not within the window
|
// Not within the window
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
l.Debug("rejected a packet (top)", "current", b.current, "incoming", i)
|
l.Debug("rejected a packet (top)",
|
||||||
|
"current", b.current,
|
||||||
|
"incoming", i,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// Update has three branches:
|
|
||||||
// - i == b.current+1: fast path; advance the cursor by one and lose-count
|
|
||||||
// the slot we just stomped (only past warmup; see the i > b.length guard
|
|
||||||
// below).
|
|
||||||
// - i > b.current+1: jump path; clear all slots between current and i
|
|
||||||
// (or up to a full window's worth, whichever is smaller) via clearRange,
|
|
||||||
// then mark i. Two arms here: a warmup arm that handles the very first
|
|
||||||
// window before the cursor has slid, and a steady-state arm that treats
|
|
||||||
// every cleared empty slot as a lost packet.
|
|
||||||
// - i <= b.current: in-window check for duplicates; out-of-window otherwise.
|
|
||||||
//
|
|
||||||
// NewBits seeds bits[0]=1 so counter 0 looks "received" — Update never
|
|
||||||
// clears that marker during warmup (clearRange skips position 0 when
|
|
||||||
// startPos=1), and once b.current >= b.length the marker is no longer
|
|
||||||
// consulted. The marker prevents a fictitious "lost" hit on the first real
|
|
||||||
// counter.
|
|
||||||
func (b *Bits) Update(l *slog.Logger, i uint64) bool {
|
func (b *Bits) Update(l *slog.Logger, i uint64) bool {
|
||||||
// Fast path: i is the next expected counter. Split out so the function
|
// If i is the next number, return true and update current.
|
||||||
// stays small and avoids paying for the slow paths' slog argument-build
|
|
||||||
// stack frame on every call. The bit read/test/write is inlined to
|
|
||||||
// touch the backing word once.
|
|
||||||
if i == b.current+1 {
|
if i == b.current+1 {
|
||||||
pos := i & b.lengthMask
|
// Check if the oldest bit was lost since we are shifting the window by 1 and occupying it with this counter
|
||||||
word := pos >> 6
|
// The very first window can only be tracked as lost once we are on the 2nd window or greater
|
||||||
mask := uint64(1) << (pos & 63)
|
if b.bits[i%b.length] == false && i > b.length {
|
||||||
w := b.bits[word]
|
|
||||||
if i > b.length && w&mask == 0 {
|
|
||||||
b.lostCounter.Inc(1)
|
b.lostCounter.Inc(1)
|
||||||
}
|
}
|
||||||
b.bits[word] = w | mask
|
b.bits[i%b.length] = true
|
||||||
b.current = i
|
b.current = i
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
return b.updateSlow(l, i)
|
|
||||||
}
|
|
||||||
|
|
||||||
// updateSlow handles jumps, in-window backfill, dupes, and out-of-window.
|
|
||||||
func (b *Bits) updateSlow(l *slog.Logger, i uint64) bool {
|
|
||||||
// If i is a jump, adjust the window, record lost, update current, and return true
|
// If i is a jump, adjust the window, record lost, update current, and return true
|
||||||
if i > b.current {
|
if i > b.current {
|
||||||
end := i
|
lost := int64(0)
|
||||||
if end > b.current+b.length {
|
// Zero out the bits between the current and the new counter value, limited by the window size,
|
||||||
end = b.current + b.length
|
// since the window is shifting
|
||||||
}
|
for n := b.current + 1; n <= min(i, b.current+b.length); n++ {
|
||||||
count := end - b.current
|
if b.bits[n%b.length] == false && n > b.length {
|
||||||
startPos := (b.current + 1) & b.lengthMask
|
|
||||||
|
|
||||||
var lost int64
|
|
||||||
if b.current >= b.length {
|
|
||||||
// Steady state: every cleared slot is past warmup, so any unset
|
|
||||||
// bit we evict is a lost packet from the previous cycle.
|
|
||||||
wasSet := b.clearRange(startPos, count)
|
|
||||||
lost = int64(count) - int64(wasSet)
|
|
||||||
} else {
|
|
||||||
// Warmup (the very first window). Some cleared slots represent
|
|
||||||
// packets <= length where eviction is not "lost" in the usual
|
|
||||||
// sense. This branch is taken at most once per connection so we
|
|
||||||
// don't bother optimizing it.
|
|
||||||
for n := b.current + 1; n <= end; n++ {
|
|
||||||
if !b.get(n) && n > b.length {
|
|
||||||
lost++
|
lost++
|
||||||
}
|
}
|
||||||
}
|
b.bits[n%b.length] = false
|
||||||
b.clearRange(startPos, count)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Anything past the new window can never be backfilled, so it's lost.
|
// Only record any skipped packets as a result of the window moving further than the window length
|
||||||
if i > b.current+b.length {
|
// Any loss within the new window will be accounted for in future calls
|
||||||
lost += int64(i - b.current - b.length)
|
lost += max(0, int64(i-b.current-b.length))
|
||||||
}
|
|
||||||
b.lostCounter.Inc(lost)
|
b.lostCounter.Inc(lost)
|
||||||
|
|
||||||
b.set(i)
|
b.bits[i%b.length] = true
|
||||||
b.current = i
|
b.current = i
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// If i is within the current window but below the current counter, check to see if it's a duplicate
|
// If i is within the current window but below the current counter,
|
||||||
if b.strictlyWithinWindow(i) {
|
// Check to see if it's a duplicate
|
||||||
pos := i & b.lengthMask
|
if i > b.current-b.length || i < b.length && b.current < b.length {
|
||||||
word := pos >> 6
|
if b.current == i || b.bits[i%b.length] == true {
|
||||||
mask := uint64(1) << (pos & 63)
|
|
||||||
w := b.bits[word]
|
|
||||||
if b.current == i || w&mask != 0 {
|
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
l.Debug("Receive window",
|
l.Debug("Receive window",
|
||||||
"accepted", false,
|
"accepted", false,
|
||||||
@@ -245,7 +104,7 @@ func (b *Bits) updateSlow(l *slog.Logger, i uint64) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
b.bits[word] = w | mask
|
b.bits[i%b.length] = true
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+128
-275
@@ -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
|
|
||||||
}
|
|
||||||
|
|||||||
+5
-30
@@ -97,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
|
||||||
@@ -184,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 {
|
||||||
@@ -283,17 +261,15 @@ 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
|
||||||
var b []byte
|
var b []byte
|
||||||
@@ -318,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)
|
||||||
}
|
}
|
||||||
@@ -329,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)
|
||||||
}
|
}
|
||||||
@@ -340,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)
|
||||||
}
|
}
|
||||||
@@ -356,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
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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,14 +57,12 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return fmt.Errorf("error while unmarshaling cert: %s", err)
|
return fmt.Errorf("error while unmarshaling cert: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !qrToStdout {
|
|
||||||
if *pf.json {
|
if *pf.json {
|
||||||
jsonCerts = append(jsonCerts, c)
|
jsonCerts = append(jsonCerts, c)
|
||||||
} else {
|
} else {
|
||||||
_, _ = out.Write([]byte(c.String()))
|
_, _ = out.Write([]byte(c.String()))
|
||||||
_, _ = out.Write([]byte("\n"))
|
_, _ = out.Write([]byte("\n"))
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
if *pf.outQRPath != "" {
|
if *pf.outQRPath != "" {
|
||||||
b, err := c.MarshalPEM()
|
b, err := c.MarshalPEM()
|
||||||
@@ -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
|
||||||
|
|||||||
+16
-38
@@ -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,11 +266,17 @@ 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 == "" {
|
||||||
|
*sf.outKeyPath = *sf.name + ".key"
|
||||||
|
}
|
||||||
|
|
||||||
|
if *sf.outCertPath == "" {
|
||||||
|
*sf.outCertPath = *sf.name + ".crt"
|
||||||
|
}
|
||||||
|
|
||||||
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
||||||
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
var crts []cert.Certificate
|
var crts []cert.Certificate
|
||||||
|
|
||||||
@@ -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`)
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -53,12 +53,7 @@ func main() {
|
|||||||
l := logging.NewLogger(os.Stdout)
|
l := logging.NewLogger(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")
|
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := doService(configPath, Build, serviceFlag); err != nil {
|
|
||||||
l.Error("Service command failed", "error", err)
|
l.Error("Service command failed", "error", err)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
@@ -66,13 +61,10 @@ 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
|
|
||||||
}
|
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
err := c.Load(*configPath)
|
err := c.Load(*configPath)
|
||||||
@@ -98,14 +90,15 @@ 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.Error("Nebula stopped due to fatal error", "error", err)
|
||||||
os.Exit(2)
|
os.Exit(2)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
"os"
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
|
||||||
"github.com/kardianos/service"
|
"github.com/kardianos/service"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
@@ -15,6 +16,7 @@ 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
|
||||||
}
|
}
|
||||||
@@ -40,47 +42,39 @@ func (p *program) Start(s service.Service) error {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
p.control, err = nebula.Main(c, false, Build, l, nil)
|
p.control, err = nebula.Main(c, *p.configTest, 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,6 +86,7 @@ func doService(configPath *string, build string, serviceFlag *string) error {
|
|||||||
|
|
||||||
prg := &program{
|
prg := &program{
|
||||||
configPath: configPath,
|
configPath: configPath,
|
||||||
|
configTest: configTest,
|
||||||
build: build,
|
build: build,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -123,9 +118,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])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+5
-7
@@ -50,13 +50,10 @@ 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 := logging.NewLogger(os.Stdout)
|
||||||
|
|
||||||
@@ -84,7 +81,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,7 +90,7 @@ 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.Error("Nebula stopped due to fatal error", "error", err)
|
||||||
os.Exit(2)
|
os.Exit(2)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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))
|
|
||||||
}
|
|
||||||
+47
-9
@@ -11,6 +11,7 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
@@ -44,6 +45,8 @@ type connectionManager struct {
|
|||||||
inactivityTimeout atomic.Int64
|
inactivityTimeout atomic.Int64
|
||||||
dropInactive atomic.Bool
|
dropInactive atomic.Bool
|
||||||
|
|
||||||
|
metricsTxPunchy metrics.Counter
|
||||||
|
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -54,6 +57,7 @@ func newConnectionManagerFromConfig(l *slog.Logger, c *config.C, hm *HostMap, p
|
|||||||
punchy: p,
|
punchy: p,
|
||||||
relayUsed: make(map[uint32]struct{}),
|
relayUsed: make(map[uint32]struct{}),
|
||||||
relayUsedLock: &sync.RWMutex{},
|
relayUsedLock: &sync.RWMutex{},
|
||||||
|
metricsTxPunchy: metrics.GetOrRegisterCounter("messages.tx.punchy", nil),
|
||||||
}
|
}
|
||||||
|
|
||||||
cm.reload(c, true)
|
cm.reload(c, true)
|
||||||
@@ -136,6 +140,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()
|
||||||
@@ -298,8 +310,8 @@ func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo
|
|||||||
} else {
|
} else {
|
||||||
cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu))
|
cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||||
cm.l.Info("send CreateRelayRequest",
|
cm.l.Info("send CreateRelayRequest",
|
||||||
"relayFrom", 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,
|
||||||
@@ -357,7 +369,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
|
|
||||||
if !outTraffic {
|
if !outTraffic {
|
||||||
// Send a punch packet to keep the NAT state alive
|
// Send a punch packet to keep the NAT state alive
|
||||||
cm.punchy.SendPunch(hostinfo)
|
cm.sendPunch(hostinfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
return decision, hostinfo, primary
|
return decision, hostinfo, primary
|
||||||
@@ -388,16 +400,17 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
|
|
||||||
// If we aren't sending or receiving traffic then its an unused tunnel and we don't to test the tunnel.
|
// If we aren't sending or receiving traffic then its an unused tunnel and we don't to test the tunnel.
|
||||||
// Just maintain NAT state if configured to do so.
|
// Just maintain NAT state if configured to do so.
|
||||||
cm.punchy.SendPunch(hostinfo)
|
cm.sendPunch(hostinfo)
|
||||||
cm.trafficTimer.Add(hostinfo.localIndexId, cm.checkInterval)
|
cm.trafficTimer.Add(hostinfo.localIndexId, cm.checkInterval)
|
||||||
return doNothing, nil, nil
|
return doNothing, nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// We aren't receiving traffic but we are sending it. The outbound
|
if cm.punchy.GetTargetEverything() {
|
||||||
// traffic itself refreshes the primary remote's NAT state; this
|
// This is similar to the old punchy behavior with a slight optimization.
|
||||||
// fans out to non-primary remotes, but only if target_all_remotes
|
// We aren't receiving traffic but we are sending it, punch on all known
|
||||||
// is configured.
|
// ips in case we need to re-prime NAT state
|
||||||
cm.punchy.SendPunchToAll(hostinfo)
|
cm.sendPunch(hostinfo)
|
||||||
|
}
|
||||||
|
|
||||||
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(cm.l).Debug("Tunnel status",
|
hostinfo.logger(cm.l).Debug("Tunnel status",
|
||||||
@@ -499,6 +512,31 @@ func (cm *connectionManager) isInvalidCertificate(now time.Time, hostinfo *HostI
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (cm *connectionManager) sendPunch(hostinfo *HostInfo) {
|
||||||
|
if !cm.punchy.GetPunch() {
|
||||||
|
// Punching is disabled
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if cm.intf.lightHouse.IsAnyLighthouseAddr(hostinfo.vpnAddrs) {
|
||||||
|
// Do not punch to lighthouses, we assume our lighthouse update interval is good enough.
|
||||||
|
// In the event the update interval is not sufficient to maintain NAT state then a publicly available lighthouse
|
||||||
|
// would lose the ability to notify us and punchy.respond would become unreliable.
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if cm.punchy.GetTargetEverything() {
|
||||||
|
hostinfo.remotes.ForEach(cm.hostMap.GetPreferredRanges(), func(addr netip.AddrPort, preferred bool) {
|
||||||
|
cm.metricsTxPunchy.Inc(1)
|
||||||
|
cm.intf.outside.WriteTo([]byte{1}, addr)
|
||||||
|
})
|
||||||
|
|
||||||
|
} else if hostinfo.remote.IsValid() {
|
||||||
|
cm.metricsTxPunchy.Inc(1)
|
||||||
|
cm.intf.outside.WriteTo([]byte{1}, hostinfo.remote)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
func (cm *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
||||||
cs := cm.intf.pki.getCertState()
|
cs := cm.intf.pki.getCertState()
|
||||||
curCrt := hostinfo.ConnectionState.myCert
|
curCrt := hostinfo.ConnectionState.myCert
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/overlay/overlaytest"
|
"github.com/slackhq/nebula/overlay/overlaytest"
|
||||||
@@ -25,7 +26,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,7 +47,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -65,7 +65,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(test.NewLogger())
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
@@ -80,6 +80,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -129,7 +130,7 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -147,7 +148,7 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(test.NewLogger())
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
@@ -162,6 +163,7 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -213,7 +215,7 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -234,7 +236,7 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
conf.Settings["tunnels"] = map[string]any{
|
conf.Settings["tunnels"] = map[string]any{
|
||||||
"drop_inactive": true,
|
"drop_inactive": true,
|
||||||
}
|
}
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
assert.True(t, nc.dropInactive.Load())
|
assert.True(t, nc.dropInactive.Load())
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
@@ -247,6 +249,7 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -339,7 +342,7 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
cs := &CertState{
|
cs := &CertState{
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{},
|
v1Cert: &dummyCert{},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -359,7 +362,7 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(test.NewLogger())
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
ifce.connectionManager = nc
|
ifce.connectionManager = nc
|
||||||
@@ -369,6 +372,7 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
ConnectionState: &ConnectionState{
|
ConnectionState: &ConnectionState{
|
||||||
myCert: &dummyCert{},
|
myCert: &dummyCert{},
|
||||||
peerCert: cachedPeerCert,
|
peerCert: cachedPeerCert,
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|||||||
+53
-71
@@ -1,49 +1,80 @@
|
|||||||
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/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 //todo I've started seeing out-of-window messages in testing?
|
||||||
|
|
||||||
type ConnectionState struct {
|
type ConnectionState struct {
|
||||||
eKey noiseutil.CipherState
|
eKey *NebulaCipherState
|
||||||
dKey noiseutil.CipherState
|
dKey *NebulaCipherState
|
||||||
|
H *noise.HandshakeState
|
||||||
myCert cert.Certificate
|
myCert cert.Certificate
|
||||||
peerCert *cert.CachedCertificate
|
peerCert *cert.CachedCertificate
|
||||||
initiator bool
|
initiator bool
|
||||||
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(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 +88,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())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
+26
-60
@@ -53,7 +53,6 @@ type Control struct {
|
|||||||
statsStart func()
|
statsStart func()
|
||||||
dnsStart func()
|
dnsStart func()
|
||||||
lighthouseStart func()
|
lighthouseStart func()
|
||||||
networkChangeStart func(rebind func())
|
|
||||||
connectionManagerStart func(context.Context)
|
connectionManagerStart func(context.Context)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -70,29 +69,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 +104,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 +114,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 +133,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 +145,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.stateLock.Lock()
|
c.CloseAllTunnels(false)
|
||||||
c.state = StateStopped
|
|
||||||
if err := c.f.Close(); err != nil {
|
if err := c.f.Close(); err != nil {
|
||||||
c.l.Error("Close interface failed", "error", err)
|
c.l.Error("Close interface failed", "error", err)
|
||||||
}
|
}
|
||||||
|
c.stateLock.Lock()
|
||||||
|
c.state = StateStopped
|
||||||
c.stateLock.Unlock()
|
c.stateLock.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Wait blocks until nebula has fully stopped, either via Stop or an internal fatal error,
|
|
||||||
// and returns the first fatal packet reader error if there was one.
|
|
||||||
// It is safe to call from multiple goroutines and at any point in the lifecycle,
|
|
||||||
// but a Wait on a Control that is never started and never stopped will block forever.
|
|
||||||
func (c *Control) Wait() error {
|
|
||||||
return c.f.wait()
|
|
||||||
}
|
|
||||||
|
|
||||||
// ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled
|
// ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled
|
||||||
func (c *Control) ShutdownBlock() {
|
func (c *Control) ShutdownBlock() {
|
||||||
sigChan := make(chan os.Signal, 1)
|
sigChan := make(chan os.Signal, 1)
|
||||||
@@ -193,20 +170,9 @@ func (c *Control) ShutdownBlock() {
|
|||||||
c.Stop()
|
c.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change.
|
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change
|
||||||
func (c *Control) RebindUDPServer() {
|
func (c *Control) RebindUDPServer() {
|
||||||
c.stateLock.Lock()
|
_ = 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()
|
||||||
@@ -339,7 +305,7 @@ func (c *Control) CloseAllTunnels(excludeLighthouses bool) (closed int) {
|
|||||||
|
|
||||||
c.l.Debug("Sending close tunnel message",
|
c.l.Debug("Sending close tunnel message",
|
||||||
"vpnAddrs", h.vpnAddrs,
|
"vpnAddrs", h.vpnAddrs,
|
||||||
"udpAddr", h.GetRemote(),
|
"udpAddr", h.remote,
|
||||||
)
|
)
|
||||||
closed++
|
closed++
|
||||||
}
|
}
|
||||||
@@ -384,7 +350,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")
|
|
||||||
}
|
|
||||||
+6
-161
@@ -1,8 +1,6 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"log/slog"
|
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"reflect"
|
"reflect"
|
||||||
@@ -11,7 +9,6 @@ import (
|
|||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
||||||
@@ -45,7 +42,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 +56,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,9 +75,7 @@ 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,
|
||||||
@@ -124,153 +119,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))
|
||||||
|
|||||||
+21
-96
@@ -11,6 +11,7 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
|
"github.com/gaissmai/bart"
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
)
|
)
|
||||||
@@ -22,10 +23,7 @@ type dnsServer struct {
|
|||||||
dnsMap4 map[string]netip.Addr
|
dnsMap4 map[string]netip.Addr
|
||||||
dnsMap6 map[string]netip.Addr
|
dnsMap6 map[string]netip.Addr
|
||||||
hostMap *HostMap
|
hostMap *HostMap
|
||||||
pki *PKI
|
myVpnAddrsTable *bart.Lite
|
||||||
|
|
||||||
// selfHost is the cached FQDN we last seeded for ourselves
|
|
||||||
selfHost string
|
|
||||||
|
|
||||||
mux *dns.ServeMux
|
mux *dns.ServeMux
|
||||||
|
|
||||||
@@ -57,14 +55,14 @@ type dnsServer struct {
|
|||||||
// they no-op when DNS isn't enabled. Each Start invocation owns a ctx-cancel
|
// they no-op when DNS isn't enabled. Each Start invocation owns a ctx-cancel
|
||||||
// watcher that tears the listener down on nebula shutdown. The returned
|
// watcher that tears the listener down on nebula shutdown. The returned
|
||||||
// pointer is always non-nil, even on error.
|
// pointer is always non-nil, even on error.
|
||||||
func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, pki *PKI, hostMap *HostMap, c *config.C) (*dnsServer, error) {
|
func newDnsServerFromConfig(ctx context.Context, l *slog.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)
|
||||||
@@ -78,7 +76,6 @@ func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, pki *PKI, hostM
|
|||||||
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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -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
|
||||||
}
|
}
|
||||||
@@ -201,14 +200,6 @@ func (d *dnsServer) Start() {
|
|||||||
close(started)
|
close(started)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Release our slot, unless a reload already replaced us, so a dead listener can't block a future Start
|
|
||||||
d.serverMu.Lock()
|
|
||||||
if d.server == server {
|
|
||||||
d.server = nil
|
|
||||||
d.started = nil
|
|
||||||
}
|
|
||||||
d.serverMu.Unlock()
|
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
d.l.Warn("Failed to run the DNS responder", "error", err)
|
d.l.Warn("Failed to run the DNS responder", "error", err)
|
||||||
}
|
}
|
||||||
@@ -258,20 +249,6 @@ func (d *dnsServer) QueryCert(data string) string {
|
|||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
// The hostmap only ever contains peers we have handshaked with, so it never carries an entry for ourselves.
|
|
||||||
// Answer self lookups straight from the local cert state.
|
|
||||||
if cs := d.certState(); cs != nil && cs.myVpnAddrsTable != nil && cs.myVpnAddrsTable.Contains(ip) {
|
|
||||||
c := cs.GetDefaultCertificate()
|
|
||||||
if c == nil {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
b, err := c.MarshalJSON()
|
|
||||||
if err != nil {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
return string(b)
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo := d.hostMap.QueryVpnAddr(ip)
|
hostinfo := d.hostMap.QueryVpnAddr(ip)
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
return ""
|
return ""
|
||||||
@@ -289,60 +266,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,12 +309,8 @@ 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) {
|
||||||
|
|||||||
+4
-295
@@ -9,10 +9,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/cert_test"
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -194,51 +191,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 +276,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 +338,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.
|
||||||
|
|||||||
+60
-240
@@ -16,7 +16,6 @@ import (
|
|||||||
"github.com/slackhq/nebula/cert_test"
|
"github.com/slackhq/nebula/cert_test"
|
||||||
"github.com/slackhq/nebula/e2e/router"
|
"github.com/slackhq/nebula/e2e/router"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -40,22 +39,11 @@ func BenchmarkHotPath(b *testing.B) {
|
|||||||
r.CancelFlowLogs()
|
r.CancelFlowLogs()
|
||||||
|
|
||||||
assertTunnel(b, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(b, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
// Pre-build the IP packet bytes once so the bench measures the data plane,
|
|
||||||
// not gopacket SerializeLayers overhead.
|
|
||||||
prebuilt := BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
|
||||||
|
|
||||||
// EnableFanIn switches the router to a 0-alloc routing path. Required
|
|
||||||
// for hot-path benchmarks; would conflict with GetFromUDP-using tests.
|
|
||||||
r.EnableFanIn()
|
|
||||||
|
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
|
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
myControl.InjectTunPacket(prebuilt)
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
// Release the TUN-side bytes back to the harness freelist; the bench
|
_ = r.RouteForAllUntilTxTun(theirControl)
|
||||||
// just confirms a packet arrived, the contents aren't inspected.
|
|
||||||
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
@@ -83,15 +71,11 @@ func BenchmarkHotPathRelay(b *testing.B) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
assertTunnel(b, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(b, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
|
|
||||||
prebuilt := BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
|
||||||
r.EnableFanIn()
|
|
||||||
|
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
|
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
myControl.InjectTunPacket(prebuilt)
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
_ = r.RouteForAllUntilTxTun(theirControl)
|
||||||
}
|
}
|
||||||
|
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
@@ -100,7 +84,6 @@ func BenchmarkHotPathRelay(b *testing.B) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGoodHandshake(t *testing.T) {
|
func TestGoodHandshake(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -113,7 +96,7 @@ func TestGoodHandshake(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||||
@@ -151,7 +134,6 @@ func TestGoodHandshake(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGoodHandshakeNoOverlap(t *testing.T) {
|
func TestGoodHandshakeNoOverlap(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "2001::69/24", nil) //look ma, cross-stack!
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "2001::69/24", nil) //look ma, cross-stack!
|
||||||
@@ -187,7 +169,6 @@ func TestGoodHandshakeNoOverlap(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestWrongResponderHandshake(t *testing.T) {
|
func TestWrongResponderHandshake(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.100/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.100/24", nil)
|
||||||
@@ -207,7 +188,7 @@ func TestWrongResponderHandshake(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
@@ -264,7 +245,6 @@ func TestWrongResponderHandshake(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.99/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.99/24", nil)
|
||||||
@@ -289,7 +269,7 @@ func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
@@ -347,7 +327,6 @@ func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestStage1Race(t *testing.T) {
|
func TestStage1Race(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// This tests ensures that two hosts handshaking with each other at the same time will allow traffic to flow
|
// This tests ensures that two hosts handshaking with each other at the same time will allow traffic to flow
|
||||||
// But will eventually collapse down to a single tunnel
|
// But will eventually collapse down to a single tunnel
|
||||||
|
|
||||||
@@ -368,8 +347,8 @@ func TestStage1Race(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake to start on both me and them")
|
t.Log("Trigger a handshake to start on both me and them")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
t.Log("Get both stage 1 handshake packets")
|
t.Log("Get both stage 1 handshake packets")
|
||||||
myHsForThem := myControl.GetFromUDP(true)
|
myHsForThem := myControl.GetFromUDP(true)
|
||||||
@@ -405,7 +384,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 +407,6 @@ func TestStage1Race(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUncleanShutdownRaceLoser(t *testing.T) {
|
func TestUncleanShutdownRaceLoser(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -446,20 +424,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 +443,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 +456,6 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUncleanShutdownRaceWinner(t *testing.T) {
|
func TestUncleanShutdownRaceWinner(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -498,7 +473,7 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
r.Log("Trigger a handshake from me to them")
|
r.Log("Trigger a handshake from me to them")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
@@ -506,13 +481,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 +494,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 +507,6 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRelays(t *testing.T) {
|
func TestRelays(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -555,7 +527,7 @@ func TestRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -564,7 +536,6 @@ func TestRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRelaysDontCareAboutIps(t *testing.T) {
|
func TestRelaysDontCareAboutIps(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "2001::9999/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "2001::9999/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -585,7 +556,7 @@ func TestRelaysDontCareAboutIps(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -594,7 +565,6 @@ func TestRelaysDontCareAboutIps(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestReestablishRelays(t *testing.T) {
|
func TestReestablishRelays(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -615,14 +585,14 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
|
||||||
t.Log("Ensure packet traversal from them to me via the relay")
|
t.Log("Ensure packet traversal from them to me via the relay")
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -632,12 +602,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 +624,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(theirControl)
|
p = r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -689,7 +659,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
t.Log("Assert the tunnel works the other way, too")
|
t.Log("Assert the tunnel works the other way, too")
|
||||||
for {
|
for {
|
||||||
t.Log("RouteForAllUntilTxTun")
|
t.Log("RouteForAllUntilTxTun")
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -725,72 +695,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 +728,8 @@ func TestStage1RaceRelays(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
||||||
|
|
||||||
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
r.Log("Wait for a packet from them to me")
|
r.Log("Wait for a packet from them to me")
|
||||||
p := r.RouteForAllUntilTxTun(myControl)
|
p := r.RouteForAllUntilTxTun(myControl)
|
||||||
@@ -838,7 +743,6 @@ func TestStage1RaceRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestStage1RaceRelays2(t *testing.T) {
|
func TestStage1RaceRelays2(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
@@ -871,8 +775,8 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
||||||
|
|
||||||
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
//r.RouteUntilAfterMsgType(myControl, header.Control, header.MessageNone)
|
//r.RouteUntilAfterMsgType(myControl, header.Control, header.MessageNone)
|
||||||
//r.RouteUntilAfterMsgType(theirControl, header.Control, header.MessageNone)
|
//r.RouteUntilAfterMsgType(theirControl, header.Control, header.MessageNone)
|
||||||
@@ -887,18 +791,18 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
|
|
||||||
t.Log("Wait until we remove extra tunnels")
|
t.Log("Wait until we remove extra tunnels")
|
||||||
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
||||||
myControl.GetHostmapIndexCount(),
|
len(myControl.GetHostmap().Indexes),
|
||||||
theirControl.GetHostmapIndexCount(),
|
len(theirControl.GetHostmap().Indexes),
|
||||||
relayControl.GetHostmapIndexCount(),
|
len(relayControl.GetHostmap().Indexes),
|
||||||
)
|
)
|
||||||
hostInfos := myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount()
|
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",
|
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
||||||
myControl.GetHostmapIndexCount(),
|
len(myControl.GetHostmap().Indexes),
|
||||||
theirControl.GetHostmapIndexCount(),
|
len(theirControl.GetHostmap().Indexes),
|
||||||
relayControl.GetHostmapIndexCount(),
|
len(relayControl.GetHostmap().Indexes),
|
||||||
)
|
)
|
||||||
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")
|
||||||
@@ -915,7 +819,6 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingRelays(t *testing.T) {
|
func TestRehandshakingRelays(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, relayConfig := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, relayConfig := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -936,7 +839,7 @@ func TestRehandshakingRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -992,24 +895,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 +922,6 @@ func TestRehandshakingRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingRelaysPrimary(t *testing.T) {
|
func TestRehandshakingRelaysPrimary(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// This test is the same as TestRehandshakingRelays but one of the terminal types is a primary swap winner
|
// This test is the same as TestRehandshakingRelays but one of the terminal types is a primary swap winner
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.128/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.128/24", m{"relay": m{"use_relays": true}})
|
||||||
@@ -1041,7 +943,7 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -1097,24 +999,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 +1026,6 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshaking(t *testing.T) {
|
func TestRehandshaking(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, myConfig := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.2/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, myConfig := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.2/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, theirConfig := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.1/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, theirConfig := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.1/24", nil)
|
||||||
@@ -1191,7 +1092,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 +1121,6 @@ func TestRehandshaking(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingLoser(t *testing.T) {
|
func TestRehandshakingLoser(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// The purpose of this test is that the race loser renews their certificate and rehandshakes. The final tunnel
|
// The purpose of this test is that the race loser renews their certificate and rehandshakes. The final tunnel
|
||||||
// Should be the one with the new certificate
|
// Should be the one with the new certificate
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -1291,7 +1191,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 +1219,6 @@ func TestRehandshakingLoser(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRaceRegression(t *testing.T) {
|
func TestRaceRegression(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// This test forces stage 1, stage 2, stage 1 to be received by me from them
|
// This test forces stage 1, stage 2, stage 1 to be received by me from them
|
||||||
// We had a bug where we were not finding the duplicate handshake and responding to the final stage 1 which
|
// We had a bug where we were not finding the duplicate handshake and responding to the final stage 1 which
|
||||||
// caused a cross-linked hostinfo
|
// caused a cross-linked hostinfo
|
||||||
@@ -1343,8 +1242,8 @@ func TestRaceRegression(t *testing.T) {
|
|||||||
//them rx stage:2 initiatorIndex=120607833 responderIndex=4209862089
|
//them rx stage:2 initiatorIndex=120607833 responderIndex=4209862089
|
||||||
|
|
||||||
t.Log("Start both handshakes")
|
t.Log("Start both handshakes")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
t.Log("Get both stage 1")
|
t.Log("Get both stage 1")
|
||||||
myStage1ForThem := myControl.GetFromUDP(true)
|
myStage1ForThem := myControl.GetFromUDP(true)
|
||||||
@@ -1380,7 +1279,6 @@ func TestRaceRegression(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestV2NonPrimaryWithLighthouse(t *testing.T) {
|
func TestV2NonPrimaryWithLighthouse(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "10.128.0.1/24, ff::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "10.128.0.1/24, ff::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
||||||
|
|
||||||
@@ -1421,7 +1319,6 @@ func TestV2NonPrimaryWithLighthouse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestV2NonPrimaryWithOffNetLighthouse(t *testing.T) {
|
func TestV2NonPrimaryWithOffNetLighthouse(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "2001::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "2001::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
||||||
|
|
||||||
@@ -1462,7 +1359,6 @@ func TestV2NonPrimaryWithOffNetLighthouse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestLighthouseUpdateOnReload(t *testing.T) {
|
func TestLighthouseUpdateOnReload(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
// Create the lighthouse
|
// Create the lighthouse
|
||||||
@@ -1538,7 +1434,6 @@ func TestLighthouseUpdateOnReload(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
unsafePrefix := "192.168.6.0/24"
|
unsafePrefix := "192.168.6.0/24"
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServerWithUdpAndUnsafeNetworks(cert.Version2, ca, caKey, "spooky", "10.128.0.2/24", netip.MustParseAddrPort("10.64.0.2:4242"), unsafePrefix, nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServerWithUdpAndUnsafeNetworks(cert.Version2, ca, caKey, "spooky", "10.128.0.2/24", netip.MustParseAddrPort("10.64.0.2:4242"), unsafePrefix, nil)
|
||||||
@@ -1560,7 +1455,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||||
@@ -1588,7 +1483,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
|||||||
assertUdpPacket(t, []byte("Hi from me"), myCachedPacket, myVpnIpNet[0].Addr(), spookyDest, 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), myCachedPacket, myVpnIpNet[0].Addr(), spookyDest, 80, 80)
|
||||||
|
|
||||||
//reply
|
//reply
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman"))
|
||||||
//wait for reply
|
//wait for reply
|
||||||
theirControl.WaitForType(1, 0, myControl)
|
theirControl.WaitForType(1, 0, myControl)
|
||||||
theirCachedPacket := myControl.GetFromTun(true)
|
theirCachedPacket := myControl.GetFromTun(true)
|
||||||
@@ -1603,78 +1498,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()
|
|
||||||
}
|
|
||||||
|
|||||||
+2
-57
@@ -294,12 +294,12 @@ func deadline(t *testing.T, seconds time.Duration) doneCb {
|
|||||||
|
|
||||||
func assertTunnel(t testing.TB, vpnIpA, vpnIpB netip.Addr, controlA, controlB *nebula.Control, r *router.R) {
|
func assertTunnel(t testing.TB, vpnIpA, vpnIpB netip.Addr, controlA, controlB *nebula.Control, r *router.R) {
|
||||||
// Send a packet from them to me
|
// Send a packet from them to me
|
||||||
controlB.InjectTunPacket(BuildTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B")))
|
controlB.InjectTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B"))
|
||||||
bPacket := r.RouteForAllUntilTxTun(controlA)
|
bPacket := r.RouteForAllUntilTxTun(controlA)
|
||||||
assertUdpPacket(t, []byte("Hi from B"), bPacket, vpnIpB, vpnIpA, 90, 80)
|
assertUdpPacket(t, []byte("Hi from B"), bPacket, vpnIpB, vpnIpA, 90, 80)
|
||||||
|
|
||||||
// And once more from me to them
|
// And once more from me to them
|
||||||
controlA.InjectTunPacket(BuildTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A")))
|
controlA.InjectTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A"))
|
||||||
aPacket := r.RouteForAllUntilTxTun(controlB)
|
aPacket := r.RouteForAllUntilTxTun(controlB)
|
||||||
assertUdpPacket(t, []byte("Hello from A"), aPacket, vpnIpA, vpnIpB, 90, 80)
|
assertUdpPacket(t, []byte("Hello from A"), aPacket, vpnIpA, vpnIpB, 90, 80)
|
||||||
}
|
}
|
||||||
@@ -408,58 +408,3 @@ func testLogLevelName() string {
|
|||||||
}
|
}
|
||||||
return "info"
|
return "info"
|
||||||
}
|
}
|
||||||
|
|
||||||
// BuildTunUDPPacket assembles an IP+UDP packet suitable for Control.InjectTunPacket.
|
|
||||||
// Using UDP here because it's a simpler protocol.
|
|
||||||
func BuildTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr netip.Addr, fromPort uint16, data []byte) []byte {
|
|
||||||
serialize := make([]gopacket.SerializableLayer, 0)
|
|
||||||
var netLayer gopacket.NetworkLayer
|
|
||||||
if toAddr.Is6() {
|
|
||||||
if !fromAddr.Is6() {
|
|
||||||
panic("Cant send ipv6 to ipv4")
|
|
||||||
}
|
|
||||||
ip := &layers.IPv6{
|
|
||||||
Version: 6,
|
|
||||||
NextHeader: layers.IPProtocolUDP,
|
|
||||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
|
||||||
DstIP: toAddr.Unmap().AsSlice(),
|
|
||||||
}
|
|
||||||
serialize = append(serialize, ip)
|
|
||||||
netLayer = ip
|
|
||||||
} else {
|
|
||||||
if !fromAddr.Is4() {
|
|
||||||
panic("Cant send ipv4 to ipv6")
|
|
||||||
}
|
|
||||||
|
|
||||||
ip := &layers.IPv4{
|
|
||||||
Version: 4,
|
|
||||||
TTL: 64,
|
|
||||||
Protocol: layers.IPProtocolUDP,
|
|
||||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
|
||||||
DstIP: toAddr.Unmap().AsSlice(),
|
|
||||||
}
|
|
||||||
serialize = append(serialize, ip)
|
|
||||||
netLayer = ip
|
|
||||||
}
|
|
||||||
|
|
||||||
udp := layers.UDP{
|
|
||||||
SrcPort: layers.UDPPort(fromPort),
|
|
||||||
DstPort: layers.UDPPort(toPort),
|
|
||||||
}
|
|
||||||
if err := udp.SetNetworkLayerForChecksum(netLayer); err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
buffer := gopacket.NewSerializeBuffer()
|
|
||||||
opt := gopacket.SerializeOptions{
|
|
||||||
ComputeChecksums: true,
|
|
||||||
FixLengths: true,
|
|
||||||
}
|
|
||||||
|
|
||||||
serialize = append(serialize, &udp, gopacket.Payload(data))
|
|
||||||
if err := gopacket.SerializeLayers(buffer, opt, serialize...); err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return buffer.Bytes()
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,47 +0,0 @@
|
|||||||
//go:build e2e_testing
|
|
||||||
// +build e2e_testing
|
|
||||||
|
|
||||||
package e2e
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/cert_test"
|
|
||||||
"github.com/slackhq/nebula/e2e/router"
|
|
||||||
"go.uber.org/goleak"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestNoGoroutineLeaks brings up two nebula instances, completes a tunnel,
|
|
||||||
// stops both, and asserts no goroutines leak past the shutdown. goleak's
|
|
||||||
// retry mechanism gives the wg.Wait()-driven goroutines a moment to drain
|
|
||||||
// before failing the assertion.
|
|
||||||
//
|
|
||||||
// Intentionally NOT t.Parallel()'d: concurrent tests would have their own
|
|
||||||
// goroutines running and trip the assertion.
|
|
||||||
func TestNoGoroutineLeaks(t *testing.T) {
|
|
||||||
defer goleak.VerifyNone(t)
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
r := router.NewR(t, myControl, theirControl)
|
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
r.RenderFlow()
|
|
||||||
|
|
||||||
// Settle period: Stop() is non-blocking; the wg-driven goroutines need
|
|
||||||
// a moment to drain. goleak retries internally too, but a short explicit
|
|
||||||
// settle reduces flakes when the suite is busy.
|
|
||||||
time.Sleep(50 * time.Millisecond)
|
|
||||||
}
|
|
||||||
@@ -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()
|
|
||||||
}
|
|
||||||
+73
-309
@@ -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)))
|
||||||
|
|
||||||
|
if len(r.ignoreFlows) > 0 {
|
||||||
var h header.H
|
var h header.H
|
||||||
var parseErr error
|
err := h.Parse(p.Data)
|
||||||
if !tun {
|
if err != nil {
|
||||||
parseErr = h.Parse(p.Data)
|
panic(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Decide before copying, the copy comes from a freelist and an ignored packet would never be released
|
|
||||||
for _, i := range r.ignoreFlows {
|
for _, i := range r.ignoreFlows {
|
||||||
if tun {
|
if !tun {
|
||||||
if i.tun.HasValue && i.tun.IsTrue {
|
if i.messageType == h.Type && i.subType == h.Subtype {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
continue
|
} else if i.tun.HasValue && i.tun.IsTrue {
|
||||||
}
|
|
||||||
|
|
||||||
// A packet we could not parse has no type to match against, so no rule can ignore it
|
|
||||||
if parseErr == nil && i.messageType == h.Type && i.subType == h.Subtype {
|
|
||||||
return nil
|
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")
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-39
@@ -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
|
||||||
@@ -405,7 +367,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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+30
-47
@@ -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
|
||||||
@@ -59,8 +59,7 @@ type Firewall struct {
|
|||||||
|
|
||||||
// assignedNetworks is a list of vpn networks assigned to us in the certificate.
|
// assignedNetworks is a list of vpn networks assigned to us in the certificate.
|
||||||
assignedNetworks []netip.Prefix
|
assignedNetworks []netip.Prefix
|
||||||
// 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
|
||||||
@@ -159,9 +158,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{
|
||||||
@@ -176,7 +176,7 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
|
|||||||
DefaultTimeout: defaultTimeout,
|
DefaultTimeout: defaultTimeout,
|
||||||
routableNetworks: routableNetworks,
|
routableNetworks: routableNetworks,
|
||||||
assignedNetworks: assignedNetworks,
|
assignedNetworks: assignedNetworks,
|
||||||
unsafeNetworks: unsafeNetworks,
|
hasUnsafeNetworks: hasUnsafeNetworks,
|
||||||
l: l,
|
l: l,
|
||||||
|
|
||||||
incomingMetrics: firewallMetrics{
|
incomingMetrics: firewallMetrics{
|
||||||
@@ -216,23 +216,23 @@ func NewFirewallFromConfig(l *slog.Logger, cs *CertState, c *config.C) (*Firewal
|
|||||||
inboundAction := c.GetString("firewall.inbound_action", "drop")
|
inboundAction := c.GetString("firewall.inbound_action", "drop")
|
||||||
switch inboundAction {
|
switch inboundAction {
|
||||||
case "reject":
|
case "reject":
|
||||||
fw.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.Warn("invalid firewall.inbound_action, defaulting to `drop`", "action", inboundAction)
|
||||||
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.Warn("invalid firewall.outbound_action, defaulting to `drop`", "action", outboundAction)
|
||||||
fw.OutboundSendReject = false
|
fw.OutSendReject = false
|
||||||
}
|
}
|
||||||
|
|
||||||
err := AddFirewallRulesFromConfig(l, false, c, fw)
|
err := AddFirewallRulesFromConfig(l, false, c, fw)
|
||||||
@@ -423,6 +423,11 @@ var ErrNoMatchingRule = errors.New("no matching rule in firewall table")
|
|||||||
// Drop returns an error if the packet should be dropped, explaining why. It
|
// Drop returns an error if the packet should be dropped, explaining why. It
|
||||||
// returns nil if the packet should not be dropped.
|
// returns nil if the packet should not be dropped.
|
||||||
func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) error {
|
func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) error {
|
||||||
|
// Check if we spoke to this tuple, if we did then allow this packet
|
||||||
|
if f.inConns(fp, h, caPool, localCache) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Make sure remote address matches nebula certificate, and determine how to treat it
|
// Make sure remote address matches nebula certificate, and determine how to treat it
|
||||||
if h.networks == nil {
|
if h.networks == nil {
|
||||||
// Simple case: Certificate has one address and no unsafe networks
|
// Simple case: Certificate has one address and no unsafe networks
|
||||||
@@ -456,11 +461,6 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *
|
|||||||
return ErrInvalidLocalIP
|
return ErrInvalidLocalIP
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check if we spoke to this tuple, if we did then allow this packet
|
|
||||||
if f.inConns(fp, h, caPool, localCache) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
table := f.OutRules
|
table := f.OutRules
|
||||||
if incoming {
|
if incoming {
|
||||||
table = f.InRules
|
table = f.InRules
|
||||||
@@ -897,7 +897,7 @@ func (flc *firewallLocalCIDR) addRule(f *Firewall, localCidr string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if localCidr == "" {
|
if localCidr == "" {
|
||||||
if len(f.unsafeNetworks) == 0 || f.defaultLocalCIDRAny {
|
if !f.hasUnsafeNetworks || f.defaultLocalCIDRAny {
|
||||||
flc.Any = true
|
flc.Any = true
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -1055,6 +1055,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 +1064,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 +1079,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)
|
|
||||||
}
|
|
||||||
|
|||||||
+1
-223
@@ -916,159 +916,6 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
|
|||||||
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
|
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_ConntrackSourceSpoofingAcrossPeers(t *testing.T) {
|
|
||||||
l := test.NewLoggerWithOutput(&bytes.Buffer{})
|
|
||||||
|
|
||||||
myVpnNetworksTable := new(bart.Lite)
|
|
||||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("192.0.2.1/24"))
|
|
||||||
|
|
||||||
owner := &dummyCert{
|
|
||||||
name: "owner",
|
|
||||||
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.1/24")},
|
|
||||||
}
|
|
||||||
|
|
||||||
victim := &cert.CachedCertificate{
|
|
||||||
Certificate: &dummyCert{
|
|
||||||
name: "victim",
|
|
||||||
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
victimHI := HostInfo{
|
|
||||||
ConnectionState: &ConnectionState{peerCert: victim},
|
|
||||||
vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")},
|
|
||||||
}
|
|
||||||
victimHI.buildNetworks(myVpnNetworksTable, victim.Certificate)
|
|
||||||
|
|
||||||
attacker := &cert.CachedCertificate{
|
|
||||||
Certificate: &dummyCert{
|
|
||||||
name: "attacker",
|
|
||||||
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.3/24")},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
attackerHI := HostInfo{
|
|
||||||
ConnectionState: &ConnectionState{peerCert: attacker},
|
|
||||||
vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.3")},
|
|
||||||
}
|
|
||||||
attackerHI.buildNetworks(myVpnNetworksTable, attacker.Certificate)
|
|
||||||
|
|
||||||
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, owner)
|
|
||||||
// Allow any inbound traffic that passes the cert / source-IP checks.
|
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", ""))
|
|
||||||
cp := cert.NewCAPool()
|
|
||||||
|
|
||||||
flow := firewall.Packet{
|
|
||||||
LocalAddr: netip.MustParseAddr("192.0.2.1"),
|
|
||||||
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
|
||||||
LocalPort: 443,
|
|
||||||
RemotePort: 55000,
|
|
||||||
Protocol: firewall.ProtoUDP,
|
|
||||||
}
|
|
||||||
|
|
||||||
require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil),
|
|
||||||
"victim's own traffic from its own overlay IP must be allowed")
|
|
||||||
|
|
||||||
unseen := flow
|
|
||||||
unseen.RemotePort = 55001
|
|
||||||
assert.Equal(t, ErrInvalidRemoteIP, fw.Drop(unseen, true, &attackerHI, cp, nil),
|
|
||||||
"sanity: attacker forging victim's source IP must be rejected when no conntrack entry exists")
|
|
||||||
|
|
||||||
got := fw.Drop(flow, true, &attackerHI, cp, nil)
|
|
||||||
t.Logf("attacker replaying victim's 4-tuple: Drop returned %v (nil == packet ALLOWED == spoof succeeded)", got)
|
|
||||||
assert.Equal(t, ErrInvalidRemoteIP, got,
|
|
||||||
"SECURITY: attacker spoofed victim's overlay source IP (192.0.2.2) by reusing an existing conntrack 4-tuple; Drop returned %v instead of rejecting", got)
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkFirewallDropConntrackHit measures Drop on an already-established flow
|
|
||||||
// (a conntrack hit). This is the fast path that the source-IP<->cert binding
|
|
||||||
// reordering adds work to, so it quantifies the cost of moving the address checks
|
|
||||||
// ahead of the conntrack lookup. Cases:
|
|
||||||
// - simple: peer cert has one address, no unsafe networks (h.networks == nil),
|
|
||||||
// so the remote-address check is a single netip.Addr compare.
|
|
||||||
// - complex: peer cert has unsafe networks (h.networks populated), so the
|
|
||||||
// remote-address check is a BART lookup.
|
|
||||||
// - noCache/localCache: whether a per-batch ConntrackCache is supplied, which in
|
|
||||||
// the original code let the fast path skip straight past the address checks.
|
|
||||||
func BenchmarkFirewallDropConntrackHit(b *testing.B) {
|
|
||||||
l := test.NewLoggerWithOutput(&bytes.Buffer{})
|
|
||||||
|
|
||||||
myVpnNetworksTable := new(bart.Lite)
|
|
||||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("192.0.2.1/24"))
|
|
||||||
|
|
||||||
owner := &dummyCert{
|
|
||||||
name: "owner",
|
|
||||||
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.1/24")},
|
|
||||||
}
|
|
||||||
|
|
||||||
simpleCert := &cert.CachedCertificate{
|
|
||||||
Certificate: &dummyCert{
|
|
||||||
name: "simple",
|
|
||||||
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
simpleHost := &HostInfo{
|
|
||||||
ConnectionState: &ConnectionState{peerCert: simpleCert},
|
|
||||||
vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")},
|
|
||||||
}
|
|
||||||
simpleHost.buildNetworks(myVpnNetworksTable, simpleCert.Certificate)
|
|
||||||
|
|
||||||
complexCert := &cert.CachedCertificate{
|
|
||||||
Certificate: &dummyCert{
|
|
||||||
name: "complex",
|
|
||||||
networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")},
|
|
||||||
unsafeNetworks: []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
complexHost := &HostInfo{
|
|
||||||
ConnectionState: &ConnectionState{peerCert: complexCert},
|
|
||||||
vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")},
|
|
||||||
}
|
|
||||||
complexHost.buildNetworks(myVpnNetworksTable, complexCert.Certificate)
|
|
||||||
|
|
||||||
cp := cert.NewCAPool()
|
|
||||||
|
|
||||||
flow := firewall.Packet{
|
|
||||||
LocalAddr: netip.MustParseAddr("192.0.2.1"),
|
|
||||||
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
|
||||||
LocalPort: 443,
|
|
||||||
RemotePort: 55000,
|
|
||||||
Protocol: firewall.ProtoUDP,
|
|
||||||
}
|
|
||||||
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
host *HostInfo
|
|
||||||
useCache bool
|
|
||||||
}{
|
|
||||||
{"simple/noCache", simpleHost, false},
|
|
||||||
{"simple/localCache", simpleHost, true},
|
|
||||||
{"complex/noCache", complexHost, false},
|
|
||||||
{"complex/localCache", complexHost, true},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range cases {
|
|
||||||
b.Run(tc.name, func(b *testing.B) {
|
|
||||||
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, owner)
|
|
||||||
require.NoError(b, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", ""))
|
|
||||||
|
|
||||||
// Establish the conntrack entry so every benchmarked Drop is a hit.
|
|
||||||
require.NoError(b, fw.Drop(flow, true, tc.host, cp, nil))
|
|
||||||
|
|
||||||
var cache firewall.ConntrackCache
|
|
||||||
if tc.useCache {
|
|
||||||
cache = firewall.ConntrackCache{}
|
|
||||||
}
|
|
||||||
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
if err := fw.Drop(flow, true, tc.host, cp, cache); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkLookup(b *testing.B) {
|
func BenchmarkLookup(b *testing.B) {
|
||||||
ml := func(m map[string]struct{}, a [][]string) {
|
ml := func(m map[string]struct{}, a [][]string) {
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
@@ -1182,80 +1029,11 @@ 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(test.NewLogger())
|
||||||
|
|||||||
@@ -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,10 +9,10 @@ require (
|
|||||||
github.com/armon/go-radix v1.0.0
|
github.com/armon/go-radix v1.0.0
|
||||||
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432
|
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432
|
||||||
github.com/flynn/noise v1.1.0
|
github.com/flynn/noise v1.1.0
|
||||||
github.com/gaissmai/bart v0.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
|
||||||
@@ -22,17 +22,16 @@ require (
|
|||||||
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
||||||
github.com/stretchr/testify v1.11.1
|
github.com/stretchr/testify v1.11.1
|
||||||
github.com/vishvananda/netlink v1.3.1
|
github.com/vishvananda/netlink v1.3.1
|
||||||
go.uber.org/goleak v1.3.0
|
|
||||||
go.yaml.in/yaml/v3 v3.0.4
|
go.yaml.in/yaml/v3 v3.0.4
|
||||||
golang.org/x/crypto v0.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 +49,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=
|
||||||
@@ -162,16 +162,16 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
|
|||||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||||
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||||
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
||||||
golang.org/x/crypto v0.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 +182,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
|
|||||||
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
||||||
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
||||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||||
golang.org/x/net v0.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 +191,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
|||||||
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.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 +208,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w
|
|||||||
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.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 +223,8 @@ golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtn
|
|||||||
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
||||||
golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE=
|
golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE=
|
||||||
golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA=
|
golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA=
|
||||||
golang.org/x/tools v0.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 +233,8 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu
|
|||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b h1:J1CaxgLerRR5lgx3wnr6L04cJFbWoceSK9JWBdglINo=
|
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b h1:J1CaxgLerRR5lgx3wnr6L04cJFbWoceSK9JWBdglINo=
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b/go.mod h1:tqur9LnfstdR9ep2LaJT4lFUl0EjlHtge+gAjmsHUG4=
|
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b/go.mod h1:tqur9LnfstdR9ep2LaJT4lFUl0EjlHtge+gAjmsHUG4=
|
||||||
golang.zx2c4.com/wireguard/windows 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
|
|
||||||
}
|
|
||||||
+813
@@ -0,0 +1,813 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"log/slog"
|
||||||
|
"net/netip"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
|
)
|
||||||
|
|
||||||
|
// NOISE IX Handshakes
|
||||||
|
|
||||||
|
// This function constructs a handshake packet, but does not actually send it
|
||||||
|
// Sending is done by the handshake manager
|
||||||
|
func ixHandshakeStage0(f *Interface, hh *HandshakeHostInfo) bool {
|
||||||
|
err := f.handshakeManager.allocateIndex(hh)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to generate index",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
cs := f.pki.getCertState()
|
||||||
|
v := cs.initiatingVersion
|
||||||
|
if hh.initiatingVersionOverride != cert.VersionPre1 {
|
||||||
|
v = hh.initiatingVersionOverride
|
||||||
|
} else if v < cert.Version2 {
|
||||||
|
// If we're connecting to a v6 address we should encourage use of a V2 cert
|
||||||
|
for _, a := range hh.hostinfo.vpnAddrs {
|
||||||
|
if a.Is6() {
|
||||||
|
v = cert.Version2
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
crt := cs.getCertificate(v)
|
||||||
|
if crt == nil {
|
||||||
|
f.l.Error("Unable to handshake with host because no certificate is available",
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
"certVersion", v,
|
||||||
|
)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
crtHs := cs.getHandshakeBytes(v)
|
||||||
|
if crtHs == nil {
|
||||||
|
f.l.Error("Unable to handshake with host because no certificate handshake bytes is available",
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
"certVersion", v,
|
||||||
|
)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
ci, err := NewConnectionState(cs, crt, true, noise.HandshakeIX)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to create connection state",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
"certVersion", v,
|
||||||
|
)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
hh.hostinfo.ConnectionState = ci
|
||||||
|
|
||||||
|
hs := &NebulaHandshake{
|
||||||
|
Details: &NebulaHandshakeDetails{
|
||||||
|
InitiatorIndex: hh.hostinfo.localIndexId,
|
||||||
|
Time: uint64(time.Now().UnixNano()),
|
||||||
|
Cert: crtHs,
|
||||||
|
CertVersion: uint32(v),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
hsBytes, err := hs.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to marshal handshake message",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"certVersion", v,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
h := header.Encode(make([]byte, header.Len), header.Version, header.Handshake, header.HandshakeIXPSK0, 0, 1)
|
||||||
|
|
||||||
|
msg, _, _, err := ci.H.WriteMessage(h, hsBytes)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to call noise.WriteMessage",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// We are sending handshake packet 1, so we don't expect to receive
|
||||||
|
// handshake packet 1 from the responder
|
||||||
|
ci.window.Update(f.l, 1)
|
||||||
|
|
||||||
|
hh.hostinfo.HandshakePacket[0] = msg
|
||||||
|
hh.ready = true
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func ixHandshakeStage1(f *Interface, via ViaSender, packet []byte, h *header.H) {
|
||||||
|
cs := f.pki.getCertState()
|
||||||
|
crt := cs.GetDefaultCertificate()
|
||||||
|
if crt == nil {
|
||||||
|
f.l.Error("Unable to handshake with host because no certificate is available",
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
"certVersion", cs.initiatingVersion,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ci, err := NewConnectionState(cs, crt, false, noise.HandshakeIX)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to create connection state",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Mark packet 1 as seen so it doesn't show up as missed
|
||||||
|
ci.window.Update(f.l, 1)
|
||||||
|
|
||||||
|
msg, _, _, err := ci.H.ReadMessage(nil, packet[header.Len:])
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to call noise.ReadMessage",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hs := &NebulaHandshake{}
|
||||||
|
err = hs.Unmarshal(msg)
|
||||||
|
if err != nil || hs.Details == nil {
|
||||||
|
f.l.Error("Failed unmarshal handshake message",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
rc, err := cert.Recombine(cert.Version(hs.Details.CertVersion), hs.Details.Cert, ci.H.PeerStatic(), ci.Curve())
|
||||||
|
if err != nil {
|
||||||
|
f.l.Info("Handshake did not contain a certificate",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
remoteCert, err := f.pki.GetCAPool().VerifyCertificate(time.Now(), rc)
|
||||||
|
if err != nil {
|
||||||
|
fp, fperr := rc.Fingerprint()
|
||||||
|
if fperr != nil {
|
||||||
|
fp = "<error generating certificate fingerprint>"
|
||||||
|
}
|
||||||
|
|
||||||
|
attrs := []slog.Attr{
|
||||||
|
slog.Any("error", err),
|
||||||
|
slog.Any("from", via),
|
||||||
|
slog.Any("handshake", m{"stage": 1, "style": "ix_psk0"}),
|
||||||
|
slog.Any("certVpnNetworks", rc.Networks()),
|
||||||
|
slog.String("certFingerprint", fp),
|
||||||
|
}
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
attrs = append(attrs, slog.Any("cert", rc))
|
||||||
|
}
|
||||||
|
|
||||||
|
// LogAttrs is intentional: attrs is a pre-built []slog.Attr slice that
|
||||||
|
// callers grow conditionally, which has no pair-form equivalent.
|
||||||
|
//nolint:sloglint
|
||||||
|
f.l.LogAttrs(context.Background(), slog.LevelInfo, "Invalid certificate from host", attrs...)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if !bytes.Equal(remoteCert.Certificate.PublicKey(), ci.H.PeerStatic()) {
|
||||||
|
f.l.Info("public key mismatch between certificate and handshake",
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
"cert", remoteCert,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if remoteCert.Certificate.Version() != ci.myCert.Version() {
|
||||||
|
// We started off using the wrong certificate version, lets see if we can match the version that was sent to us
|
||||||
|
myCertOtherVersion := cs.getCertificate(remoteCert.Certificate.Version())
|
||||||
|
if myCertOtherVersion == nil {
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("Might be unable to handshake with host due to missing certificate version",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
"cert", remoteCert,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Record the certificate we are actually using
|
||||||
|
ci.myCert = myCertOtherVersion
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(remoteCert.Certificate.Networks()) == 0 {
|
||||||
|
f.l.Info("No networks in certificate",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"cert", remoteCert,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
certName := remoteCert.Certificate.Name()
|
||||||
|
certVersion := remoteCert.Certificate.Version()
|
||||||
|
fingerprint := remoteCert.Fingerprint
|
||||||
|
issuer := remoteCert.Certificate.Issuer()
|
||||||
|
vpnNetworks := remoteCert.Certificate.Networks()
|
||||||
|
|
||||||
|
anyVpnAddrsInCommon := false
|
||||||
|
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
||||||
|
for i, network := range vpnNetworks {
|
||||||
|
if f.myVpnAddrsTable.Contains(network.Addr()) {
|
||||||
|
f.l.Error("Refusing to handshake with myself",
|
||||||
|
"vpnNetworks", vpnNetworks,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
vpnAddrs[i] = network.Addr()
|
||||||
|
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
||||||
|
anyVpnAddrsInCommon = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !via.IsRelayed {
|
||||||
|
// We only want to apply the remote allow list for direct tunnels here
|
||||||
|
if !f.lightHouse.GetRemoteAllowList().AllowAll(vpnAddrs, via.UdpAddr.Addr()) {
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
myIndex, err := generateIndex(f.l)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to generate index",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo := &HostInfo{
|
||||||
|
ConnectionState: ci,
|
||||||
|
localIndexId: myIndex,
|
||||||
|
remoteIndexId: hs.Details.InitiatorIndex,
|
||||||
|
vpnAddrs: vpnAddrs,
|
||||||
|
HandshakePacket: make(map[uint8][]byte, 0),
|
||||||
|
lastHandshakeTime: hs.Details.Time,
|
||||||
|
relayState: RelayState{
|
||||||
|
relays: nil,
|
||||||
|
relayForByAddr: map[netip.Addr]*Relay{},
|
||||||
|
relayForByIdx: map[uint32]*Relay{},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgRxL := f.l.With(
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"initiatorIndex", hs.Details.InitiatorIndex,
|
||||||
|
"responderIndex", hs.Details.ResponderIndex,
|
||||||
|
"remoteIndex", h.RemoteIndex,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
|
||||||
|
if anyVpnAddrsInCommon {
|
||||||
|
msgRxL.Info("Handshake message received")
|
||||||
|
} else {
|
||||||
|
//todo warn if not lighthouse or relay?
|
||||||
|
msgRxL.Info("Handshake message received, but no vpnNetworks in common.")
|
||||||
|
}
|
||||||
|
|
||||||
|
hs.Details.ResponderIndex = myIndex
|
||||||
|
hs.Details.Cert = cs.getHandshakeBytes(ci.myCert.Version())
|
||||||
|
if hs.Details.Cert == nil {
|
||||||
|
msgRxL.Error("Unable to handshake with host because no certificate handshake bytes is available",
|
||||||
|
"myCertVersion", ci.myCert.Version(),
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hs.Details.CertVersion = uint32(ci.myCert.Version())
|
||||||
|
// Update the time in case their clock is way off from ours
|
||||||
|
hs.Details.Time = uint64(time.Now().UnixNano())
|
||||||
|
|
||||||
|
hsBytes, err := hs.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to marshal handshake message",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
nh := header.Encode(make([]byte, header.Len), header.Version, header.Handshake, header.HandshakeIXPSK0, hs.Details.InitiatorIndex, 2)
|
||||||
|
msg, dKey, eKey, err := ci.H.WriteMessage(nh, hsBytes)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to call noise.WriteMessage",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
} else if dKey == nil || eKey == nil {
|
||||||
|
f.l.Error("Noise did not arrive at a key",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo.HandshakePacket[0] = make([]byte, len(packet[header.Len:]))
|
||||||
|
copy(hostinfo.HandshakePacket[0], packet[header.Len:])
|
||||||
|
|
||||||
|
// Regardless of whether you are the sender or receiver, you should arrive here
|
||||||
|
// and complete standing up the connection.
|
||||||
|
hostinfo.HandshakePacket[2] = make([]byte, len(msg))
|
||||||
|
copy(hostinfo.HandshakePacket[2], msg)
|
||||||
|
|
||||||
|
// We are sending handshake packet 2, so we don't expect to receive
|
||||||
|
// handshake packet 2 from the initiator.
|
||||||
|
ci.window.Update(f.l, 2)
|
||||||
|
|
||||||
|
ci.peerCert = remoteCert
|
||||||
|
ci.dKey = NewNebulaCipherState(dKey)
|
||||||
|
ci.eKey = NewNebulaCipherState(eKey)
|
||||||
|
|
||||||
|
hostinfo.remotes = f.lightHouse.QueryCache(vpnAddrs)
|
||||||
|
if !via.IsRelayed {
|
||||||
|
hostinfo.SetRemote(via.UdpAddr)
|
||||||
|
}
|
||||||
|
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
||||||
|
|
||||||
|
existing, err := f.handshakeManager.CheckAndComplete(hostinfo, 0, f)
|
||||||
|
if err != nil {
|
||||||
|
switch err {
|
||||||
|
case ErrAlreadySeen:
|
||||||
|
// Update remote if preferred
|
||||||
|
if existing.SetRemoteIfPreferred(f.hostMap, via) {
|
||||||
|
// Send a test packet to ensure the other side has also switched to
|
||||||
|
// the preferred remote
|
||||||
|
f.SendMessageToVpnAddr(header.Test, header.TestRequest, vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
||||||
|
}
|
||||||
|
|
||||||
|
msg = existing.HandshakePacket[2]
|
||||||
|
f.messageMetrics.Tx(header.Handshake, header.MessageSubType(msg[1]), 1)
|
||||||
|
if !via.IsRelayed {
|
||||||
|
err := f.outside.WriteTo(msg, via.UdpAddr)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to send handshake message",
|
||||||
|
"vpnAddrs", existing.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
"cached", true,
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
f.l.Info("Handshake message sent",
|
||||||
|
"vpnAddrs", existing.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
"cached", true,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
} else {
|
||||||
|
if via.relay == nil {
|
||||||
|
f.l.Error("Handshake send failed: both addr and via.relay are nil.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||||
|
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||||
|
f.l.Info("Handshake message sent",
|
||||||
|
"vpnAddrs", existing.vpnAddrs,
|
||||||
|
"relay", via.relayHI.vpnAddrs[0],
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
"cached", true,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
case ErrExistingHostInfo:
|
||||||
|
// This means there was an existing tunnel and this handshake was older than the one we are currently based on
|
||||||
|
f.l.Info("Handshake too old",
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"oldHandshakeTime", existing.lastHandshakeTime,
|
||||||
|
"newHandshakeTime", hostinfo.lastHandshakeTime,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"initiatorIndex", hs.Details.InitiatorIndex,
|
||||||
|
"responderIndex", hs.Details.ResponderIndex,
|
||||||
|
"remoteIndex", h.RemoteIndex,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
|
||||||
|
// Send a test packet to trigger an authenticated tunnel test, this should suss out any lingering tunnel issues
|
||||||
|
f.SendMessageToVpnAddr(header.Test, header.TestRequest, vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
||||||
|
return
|
||||||
|
case ErrLocalIndexCollision:
|
||||||
|
// This means we failed to insert because of collision on localIndexId. Just let the next handshake packet retry
|
||||||
|
f.l.Error("Failed to add HostInfo due to localIndex collision",
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"initiatorIndex", hs.Details.InitiatorIndex,
|
||||||
|
"responderIndex", hs.Details.ResponderIndex,
|
||||||
|
"remoteIndex", h.RemoteIndex,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
"localIndex", hostinfo.localIndexId,
|
||||||
|
"collision", existing.vpnAddrs,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
// Shouldn't happen, but just in case someone adds a new error type to CheckAndComplete
|
||||||
|
// And we forget to update it here
|
||||||
|
f.l.Error("Failed to add HostInfo to HostMap",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"initiatorIndex", hs.Details.InitiatorIndex,
|
||||||
|
"responderIndex", hs.Details.ResponderIndex,
|
||||||
|
"remoteIndex", h.RemoteIndex,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Do the send
|
||||||
|
f.messageMetrics.Tx(header.Handshake, header.MessageSubType(msg[1]), 1)
|
||||||
|
if !via.IsRelayed {
|
||||||
|
err = f.outside.WriteTo(msg, via.UdpAddr)
|
||||||
|
log := f.l.With(
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"initiatorIndex", hs.Details.InitiatorIndex,
|
||||||
|
"responderIndex", hs.Details.ResponderIndex,
|
||||||
|
"remoteIndex", h.RemoteIndex,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to send handshake", "error", err)
|
||||||
|
} else {
|
||||||
|
log.Info("Handshake message sent")
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if via.relay == nil {
|
||||||
|
f.l.Error("Handshake send failed: both addr and via.relay are nil.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||||
|
// I successfully received a handshake. Just in case I marked this tunnel as 'Disestablished', ensure
|
||||||
|
// it's correctly marked as working.
|
||||||
|
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
|
||||||
|
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||||
|
f.l.Info("Handshake message sent",
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"relay", via.relayHI.vpnAddrs[0],
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"initiatorIndex", hs.Details.InitiatorIndex,
|
||||||
|
"responderIndex", hs.Details.ResponderIndex,
|
||||||
|
"remoteIndex", h.RemoteIndex,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
f.connectionManager.AddTrafficWatch(hostinfo)
|
||||||
|
|
||||||
|
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||||
|
|
||||||
|
// Don't wait for UpdateWorker
|
||||||
|
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
||||||
|
f.lightHouse.TriggerUpdate()
|
||||||
|
}
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
func ixHandshakeStage2(f *Interface, via ViaSender, hh *HandshakeHostInfo, packet []byte, h *header.H) bool {
|
||||||
|
if hh == nil {
|
||||||
|
// Nothing here to tear down, got a bogus stage 2 packet
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
hh.Lock()
|
||||||
|
defer hh.Unlock()
|
||||||
|
|
||||||
|
hostinfo := hh.hostinfo
|
||||||
|
if !via.IsRelayed {
|
||||||
|
// The vpnAddr we know about is the one we tried to handshake with, use it to apply the remote allow list.
|
||||||
|
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
ci := hostinfo.ConnectionState
|
||||||
|
msg, eKey, dKey, err := ci.H.ReadMessage(nil, packet[header.Len:])
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to call noise.ReadMessage",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
"header", h,
|
||||||
|
)
|
||||||
|
|
||||||
|
// We don't want to tear down the connection on a bad ReadMessage because it could be an attacker trying
|
||||||
|
// to DOS us. Every other error condition after should to allow a possible good handshake to complete in the
|
||||||
|
// near future
|
||||||
|
return false
|
||||||
|
} else if dKey == nil || eKey == nil {
|
||||||
|
f.l.Error("Noise did not arrive at a key",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
|
||||||
|
// This should be impossible in IX but just in case, if we get here then there is no chance to recover
|
||||||
|
// the handshake state machine. Tear it down
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
hs := &NebulaHandshake{}
|
||||||
|
err = hs.Unmarshal(msg)
|
||||||
|
if err != nil || hs.Details == nil {
|
||||||
|
f.l.Error("Failed unmarshal handshake message",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
|
||||||
|
// The handshake state machine is complete, if things break now there is no chance to recover. Tear down and start again
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
rc, err := cert.Recombine(cert.Version(hs.Details.CertVersion), hs.Details.Cert, ci.H.PeerStatic(), ci.Curve())
|
||||||
|
if err != nil {
|
||||||
|
f.l.Info("Handshake did not contain a certificate",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
remoteCert, err := f.pki.GetCAPool().VerifyCertificate(time.Now(), rc)
|
||||||
|
if err != nil {
|
||||||
|
fp, err := rc.Fingerprint()
|
||||||
|
if err != nil {
|
||||||
|
fp = "<error generating certificate fingerprint>"
|
||||||
|
}
|
||||||
|
|
||||||
|
attrs := []slog.Attr{
|
||||||
|
slog.Any("error", err),
|
||||||
|
slog.Any("from", via),
|
||||||
|
slog.Any("vpnAddrs", hostinfo.vpnAddrs),
|
||||||
|
slog.Any("handshake", m{"stage": 2, "style": "ix_psk0"}),
|
||||||
|
slog.String("certFingerprint", fp),
|
||||||
|
slog.Any("certVpnNetworks", rc.Networks()),
|
||||||
|
}
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
attrs = append(attrs, slog.Any("cert", rc))
|
||||||
|
}
|
||||||
|
|
||||||
|
// LogAttrs is intentional: attrs is a pre-built []slog.Attr slice that
|
||||||
|
// callers grow conditionally, which has no pair-form equivalent.
|
||||||
|
//nolint:sloglint
|
||||||
|
f.l.LogAttrs(context.Background(), slog.LevelInfo, "Invalid certificate from host", attrs...)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if !bytes.Equal(remoteCert.Certificate.PublicKey(), ci.H.PeerStatic()) {
|
||||||
|
f.l.Info("public key mismatch between certificate and handshake",
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
"cert", remoteCert,
|
||||||
|
)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(remoteCert.Certificate.Networks()) == 0 {
|
||||||
|
f.l.Info("No networks in certificate",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"cert", remoteCert,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
vpnNetworks := remoteCert.Certificate.Networks()
|
||||||
|
certName := remoteCert.Certificate.Name()
|
||||||
|
certVersion := remoteCert.Certificate.Version()
|
||||||
|
fingerprint := remoteCert.Fingerprint
|
||||||
|
issuer := remoteCert.Certificate.Issuer()
|
||||||
|
|
||||||
|
hostinfo.remoteIndexId = hs.Details.ResponderIndex
|
||||||
|
hostinfo.lastHandshakeTime = hs.Details.Time
|
||||||
|
|
||||||
|
// Store their cert and our symmetric keys
|
||||||
|
ci.peerCert = remoteCert
|
||||||
|
ci.dKey = NewNebulaCipherState(dKey)
|
||||||
|
ci.eKey = NewNebulaCipherState(eKey)
|
||||||
|
|
||||||
|
// Make sure the current udpAddr being used is set for responding
|
||||||
|
if !via.IsRelayed {
|
||||||
|
hostinfo.SetRemote(via.UdpAddr)
|
||||||
|
} else {
|
||||||
|
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
correctHostResponded := false
|
||||||
|
anyVpnAddrsInCommon := false
|
||||||
|
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
||||||
|
for i, network := range vpnNetworks {
|
||||||
|
vpnAddrs[i] = network.Addr()
|
||||||
|
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
||||||
|
anyVpnAddrsInCommon = true
|
||||||
|
}
|
||||||
|
if hostinfo.vpnAddrs[0] == network.Addr() {
|
||||||
|
// todo is it more correct to see if any of hostinfo.vpnAddrs are in the cert? it should have len==1, but one day it might not?
|
||||||
|
correctHostResponded = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure the right host responded
|
||||||
|
if !correctHostResponded {
|
||||||
|
f.l.Info("Incorrect host responded to handshake",
|
||||||
|
"intendedVpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"haveVpnNetworks", vpnNetworks,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
|
||||||
|
// Release our old handshake from pending, it should not continue
|
||||||
|
f.handshakeManager.DeleteHostInfo(hostinfo)
|
||||||
|
|
||||||
|
// Create a new hostinfo/handshake for the intended vpn ip
|
||||||
|
//TODO is hostinfo.vpnAddrs[0] always the address to use?
|
||||||
|
f.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], func(newHH *HandshakeHostInfo) {
|
||||||
|
// Block the current used address
|
||||||
|
newHH.hostinfo.remotes = hostinfo.remotes
|
||||||
|
newHH.hostinfo.remotes.BlockRemote(via)
|
||||||
|
|
||||||
|
f.l.Info("Blocked addresses for handshakes",
|
||||||
|
"blockedUdpAddrs", newHH.hostinfo.remotes.CopyBlockedRemotes(),
|
||||||
|
"vpnNetworks", vpnNetworks,
|
||||||
|
"remotes", newHH.hostinfo.remotes.CopyAddrs(f.hostMap.GetPreferredRanges()),
|
||||||
|
)
|
||||||
|
|
||||||
|
// Swap the packet store to benefit the original intended recipient
|
||||||
|
newHH.packetStore = hh.packetStore
|
||||||
|
hh.packetStore = []*cachedPacket{}
|
||||||
|
|
||||||
|
// Finally, put the correct vpn addrs in the host info, tell them to close the tunnel, and return true to tear down
|
||||||
|
hostinfo.vpnAddrs = vpnAddrs
|
||||||
|
f.sendCloseTunnel(hostinfo)
|
||||||
|
})
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Mark packet 2 as seen so it doesn't show up as missed
|
||||||
|
ci.window.Update(f.l, 2)
|
||||||
|
|
||||||
|
duration := time.Since(hh.startTime).Nanoseconds()
|
||||||
|
msgRxL := f.l.With(
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"initiatorIndex", hs.Details.InitiatorIndex,
|
||||||
|
"responderIndex", hs.Details.ResponderIndex,
|
||||||
|
"remoteIndex", h.RemoteIndex,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
"durationNs", duration,
|
||||||
|
"sentCachedPackets", len(hh.packetStore),
|
||||||
|
)
|
||||||
|
if anyVpnAddrsInCommon {
|
||||||
|
msgRxL.Info("Handshake message received")
|
||||||
|
} else {
|
||||||
|
//todo warn if not lighthouse or relay?
|
||||||
|
msgRxL.Info("Handshake message received, but no vpnNetworks in common.")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build up the radix for the firewall if we have subnets in the cert
|
||||||
|
hostinfo.vpnAddrs = vpnAddrs
|
||||||
|
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
||||||
|
|
||||||
|
// Complete our handshake and update metrics, this will replace any existing tunnels for the vpnAddrs here
|
||||||
|
f.handshakeManager.Complete(hostinfo, f)
|
||||||
|
f.connectionManager.AddTrafficWatch(hostinfo)
|
||||||
|
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Sending stored packets",
|
||||||
|
"count", len(hh.packetStore),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(hh.packetStore) > 0 {
|
||||||
|
nb := make([]byte, 12, 12)
|
||||||
|
out := make([]byte, mtu)
|
||||||
|
for _, cp := range hh.packetStore {
|
||||||
|
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
||||||
|
}
|
||||||
|
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||||
|
f.metricHandshakes.Update(duration)
|
||||||
|
|
||||||
|
// Don't wait for UpdateWorker
|
||||||
|
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
||||||
|
f.lightHouse.TriggerUpdate()
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
+178
-614
@@ -14,7 +14,6 @@ import (
|
|||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/handshake"
|
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
@@ -23,18 +22,7 @@ const (
|
|||||||
DefaultHandshakeTryInterval = time.Millisecond * 100
|
DefaultHandshakeTryInterval = time.Millisecond * 100
|
||||||
DefaultHandshakeRetries = 10
|
DefaultHandshakeRetries = 10
|
||||||
DefaultHandshakeTriggerBuffer = 64
|
DefaultHandshakeTriggerBuffer = 64
|
||||||
|
DefaultUseRelays = true
|
||||||
// maxCachedPackets is how many unsent packets we'll buffer per pending
|
|
||||||
// handshake before dropping further ones.
|
|
||||||
maxCachedPackets = 100
|
|
||||||
|
|
||||||
// HandshakePacket map keys mirror the IX protocol stage convention:
|
|
||||||
// stage 0 = the initiator's first message (and what the responder
|
|
||||||
// receives, stripped of header)
|
|
||||||
// stage 2 = the responder's reply
|
|
||||||
// Other handshake patterns will need new keys when added.
|
|
||||||
handshakePacketStage0 uint8 = 0
|
|
||||||
handshakePacketStage2 uint8 = 2
|
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
@@ -42,6 +30,7 @@ var (
|
|||||||
tryInterval: DefaultHandshakeTryInterval,
|
tryInterval: DefaultHandshakeTryInterval,
|
||||||
retries: DefaultHandshakeRetries,
|
retries: DefaultHandshakeRetries,
|
||||||
triggerBuffer: DefaultHandshakeTriggerBuffer,
|
triggerBuffer: DefaultHandshakeTriggerBuffer,
|
||||||
|
useRelays: DefaultUseRelays,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -49,6 +38,7 @@ type HandshakeConfig struct {
|
|||||||
tryInterval time.Duration
|
tryInterval time.Duration
|
||||||
retries int64
|
retries int64
|
||||||
triggerBuffer int
|
triggerBuffer int
|
||||||
|
useRelays bool
|
||||||
|
|
||||||
messageMetrics *MessageMetrics
|
messageMetrics *MessageMetrics
|
||||||
}
|
}
|
||||||
@@ -83,15 +73,13 @@ type HandshakeHostInfo struct {
|
|||||||
initiatingVersionOverride cert.Version // Should we use a non-default cert version for this handshake?
|
initiatingVersionOverride cert.Version // Should we use a non-default cert version for this handshake?
|
||||||
counter int64 // How many attempts have we made so far
|
counter int64 // How many attempts have we made so far
|
||||||
lastRemotes []netip.AddrPort // Remotes that we sent to during the previous attempt
|
lastRemotes []netip.AddrPort // Remotes that we sent to during the previous attempt
|
||||||
lastRelays []netip.Addr // Relays we attempted to use during the previous attempt
|
|
||||||
packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes
|
packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes
|
||||||
|
|
||||||
hostinfo *HostInfo
|
hostinfo *HostInfo
|
||||||
machine *handshake.Machine // The handshake state machine, set during stage 0 (initiator) or beginHandshake (responder multi-message)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hh *HandshakeHostInfo) cachePacket(l *slog.Logger, t header.MessageType, st header.MessageSubType, packet []byte, f packetCallback, m *cachedPacketMetrics) {
|
func (hh *HandshakeHostInfo) cachePacket(l *slog.Logger, t header.MessageType, st header.MessageSubType, packet []byte, f packetCallback, m *cachedPacketMetrics) {
|
||||||
if len(hh.packetStore) < maxCachedPackets {
|
if len(hh.packetStore) < 100 {
|
||||||
tempPacket := make([]byte, len(packet))
|
tempPacket := make([]byte, len(packet))
|
||||||
copy(tempPacket, packet)
|
copy(tempPacket, packet)
|
||||||
|
|
||||||
@@ -149,18 +137,6 @@ func (hm *HandshakeManager) Run(ctx context.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HandshakeManager) HandleIncoming(via ViaSender, packet []byte, h *header.H) {
|
func (hm *HandshakeManager) HandleIncoming(via ViaSender, packet []byte, h *header.H) {
|
||||||
// Gate on known handshake subtypes. Unknown subtypes (or future ones we
|
|
||||||
// don't yet support) are dropped here rather than silently routed through
|
|
||||||
// the IX path. Add a case when introducing a new pattern.
|
|
||||||
switch h.Subtype {
|
|
||||||
case header.HandshakeIXPSK0:
|
|
||||||
// supported
|
|
||||||
default:
|
|
||||||
hm.l.Debug("dropping handshake with unsupported subtype",
|
|
||||||
"from", via, "subtype", h.Subtype)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// First remote allow list check before we know the vpnIp
|
// First remote allow list check before we know the vpnIp
|
||||||
if !via.IsRelayed {
|
if !via.IsRelayed {
|
||||||
if !hm.lightHouse.GetRemoteAllowList().AllowUnknownVpnAddr(via.UdpAddr.Addr()) {
|
if !hm.lightHouse.GetRemoteAllowList().AllowUnknownVpnAddr(via.UdpAddr.Addr()) {
|
||||||
@@ -169,27 +145,19 @@ func (hm *HandshakeManager) HandleIncoming(via ViaSender, packet []byte, h *head
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// First message of a new handshake. The wire format requires RemoteIndex
|
switch h.Subtype {
|
||||||
// to be zero here (the initiator has no responder index to fill in yet),
|
case header.HandshakeIXPSK0:
|
||||||
// and generateIndex never allocates 0, so any non-zero RemoteIndex on a
|
switch h.MessageCounter {
|
||||||
// stage-1 packet is malformed or someone probing for an index collision.
|
case 1:
|
||||||
// Drop without paying the cost of running noise on a pending Machine.
|
ixHandshakeStage1(hm.f, via, packet, h)
|
||||||
if h.MessageCounter == 1 {
|
|
||||||
if h.RemoteIndex != 0 {
|
|
||||||
hm.l.Debug("dropping stage-1 handshake with non-zero RemoteIndex",
|
|
||||||
"from", via, "remoteIndex", h.RemoteIndex)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
hm.beginHandshake(via, packet, h)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Continuation message must match a pending handshake by index.
|
case 2:
|
||||||
// Anything else is an orphaned packet (e.g., late retransmit after
|
newHostinfo := hm.queryIndex(h.RemoteIndex)
|
||||||
// timeout) and is dropped.
|
tearDown := ixHandshakeStage2(hm.f, via, newHostinfo, packet, h)
|
||||||
if hh := hm.queryIndex(h.RemoteIndex); hh != nil {
|
if tearDown && newHostinfo != nil {
|
||||||
hm.continueHandshake(via, hh, packet)
|
hm.DeleteHostInfo(newHostinfo.hostinfo)
|
||||||
return
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -215,21 +183,13 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
hostinfo := hh.hostinfo
|
hostinfo := hh.hostinfo
|
||||||
// If we are out of time, clean up
|
// If we are out of time, clean up
|
||||||
if hh.counter >= hm.config.retries {
|
if hh.counter >= hm.config.retries {
|
||||||
fields := []any{
|
hh.hostinfo.logger(hm.l).Info("Handshake timed out",
|
||||||
"udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()),
|
"udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()),
|
||||||
"initiatorIndex", hh.hostinfo.localIndexId,
|
"initiatorIndex", hh.hostinfo.localIndexId,
|
||||||
|
"remoteIndex", hh.hostinfo.remoteIndexId,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
"durationNs", time.Since(hh.startTime).Nanoseconds(),
|
"durationNs", time.Since(hh.startTime).Nanoseconds(),
|
||||||
}
|
)
|
||||||
// hh.machine can be nil here if buildStage0Packet never succeeded
|
|
||||||
// (e.g., no certificate available). In that case there's no useful
|
|
||||||
// handshake metadata to log.
|
|
||||||
if hh.machine != nil {
|
|
||||||
fields = append(fields, "handshake", m{
|
|
||||||
"stage": uint64(hh.machine.MessageIndex()),
|
|
||||||
"style": header.SubTypeName(header.Handshake, hh.machine.Subtype()),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
hh.hostinfo.logger(hm.l).Info("Handshake timed out", fields...)
|
|
||||||
hm.metricTimedOut.Inc(1)
|
hm.metricTimedOut.Inc(1)
|
||||||
hm.DeleteHostInfo(hostinfo)
|
hm.DeleteHostInfo(hostinfo)
|
||||||
return
|
return
|
||||||
@@ -240,25 +200,12 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
|
|
||||||
// Check if we have a handshake packet to transmit yet
|
// Check if we have a handshake packet to transmit yet
|
||||||
if !hh.ready {
|
if !hh.ready {
|
||||||
if !hm.buildStage0Packet(hh) {
|
if !ixHandshakeStage0(hm.f, hh) {
|
||||||
hm.OutboundHandshakeTimer.Add(vpnIp, hm.config.tryInterval*time.Duration(hh.counter))
|
hm.OutboundHandshakeTimer.Add(vpnIp, hm.config.tryInterval*time.Duration(hh.counter))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TODO: this hardcodes "always retransmit stage 0", which is correct for
|
|
||||||
// IX (the initiator only ever sends one packet, msg1) but wrong the
|
|
||||||
// moment a 3+ message pattern lands. The retry loop should resend the
|
|
||||||
// most recent outgoing message, not always stage 0. That implies
|
|
||||||
// HandshakeHostInfo tracking a single "currentOutbound" packet (bytes +
|
|
||||||
// header metadata) that gets replaced as the handshake progresses,
|
|
||||||
// instead of indexing into HandshakePacket.
|
|
||||||
stage0 := hostinfo.HandshakePacket[handshakePacketStage0]
|
|
||||||
hsFields := m{
|
|
||||||
"stage": uint64(hh.machine.MessageIndex()),
|
|
||||||
"style": header.SubTypeName(header.Handshake, hh.machine.Subtype()),
|
|
||||||
}
|
|
||||||
|
|
||||||
// Get a remotes object if we don't already have one.
|
// Get a remotes object if we don't already have one.
|
||||||
// This is mainly to protect us as this should never be the case
|
// This is mainly to protect us as this should never be the case
|
||||||
// NB ^ This comment doesn't jive. It's how the thing gets initialized.
|
// NB ^ This comment doesn't jive. It's how the thing gets initialized.
|
||||||
@@ -292,19 +239,13 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
// Send the handshake to all known ips, stage 2 takes care of assigning the hostinfo.remote based on the first to reply
|
// Send the handshake to all known ips, stage 2 takes care of assigning the hostinfo.remote based on the first to reply
|
||||||
var sentTo []netip.AddrPort
|
var sentTo []netip.AddrPort
|
||||||
hostinfo.remotes.ForEach(hm.mainHostMap.GetPreferredRanges(), func(addr netip.AddrPort, _ bool) {
|
hostinfo.remotes.ForEach(hm.mainHostMap.GetPreferredRanges(), func(addr netip.AddrPort, _ bool) {
|
||||||
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
|
hm.messageMetrics.Tx(header.Handshake, header.MessageSubType(hostinfo.HandshakePacket[0][1]), 1)
|
||||||
err := hm.outside.WriteTo(stage0, addr)
|
err := hm.outside.WriteTo(hostinfo.HandshakePacket[0], addr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// These repeat every attempt, so match the success log below and only shout when the remotes changed
|
hostinfo.logger(hm.l).Error("Failed to send handshake message",
|
||||||
level := slog.LevelDebug
|
|
||||||
if remotesHaveChanged {
|
|
||||||
level = slog.LevelError
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo.logger(hm.l).Log(context.Background(), level, "Failed to send handshake message",
|
|
||||||
"udpAddr", addr,
|
"udpAddr", addr,
|
||||||
"initiatorIndex", hostinfo.localIndexId,
|
"initiatorIndex", hostinfo.localIndexId,
|
||||||
"handshake", hsFields,
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
"error", err,
|
"error", err,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -319,17 +260,156 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
hostinfo.logger(hm.l).Info("Handshake message sent",
|
hostinfo.logger(hm.l).Info("Handshake message sent",
|
||||||
"udpAddrs", sentTo,
|
"udpAddrs", sentTo,
|
||||||
"initiatorIndex", hostinfo.localIndexId,
|
"initiatorIndex", hostinfo.localIndexId,
|
||||||
"handshake", hsFields,
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
)
|
)
|
||||||
} else if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
} else if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(hm.l).Debug("Handshake message sent",
|
hostinfo.logger(hm.l).Debug("Handshake message sent",
|
||||||
"udpAddrs", sentTo,
|
"udpAddrs", sentTo,
|
||||||
"initiatorIndex", hostinfo.localIndexId,
|
"initiatorIndex", hostinfo.localIndexId,
|
||||||
"handshake", hsFields,
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
hm.f.relayManager.StartRelays(hm.f, vpnIp, hh, stage0)
|
if hm.config.useRelays && len(hostinfo.remotes.relays) > 0 {
|
||||||
|
hostinfo.logger(hm.l).Info("Attempt to relay through hosts", "relays", hostinfo.remotes.relays)
|
||||||
|
// Send a RelayRequest to all known Relay IP's
|
||||||
|
for _, relay := range hostinfo.remotes.relays {
|
||||||
|
// Don't relay through the host I'm trying to connect to
|
||||||
|
if relay == vpnIp {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Don't relay to myself
|
||||||
|
if hm.f.myVpnAddrsTable.Contains(relay) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
relayHostInfo := hm.mainHostMap.QueryVpnAddr(relay)
|
||||||
|
if relayHostInfo == nil || !relayHostInfo.remote.IsValid() {
|
||||||
|
hostinfo.logger(hm.l).Info("Establish tunnel to relay target", "relay", relay.String())
|
||||||
|
hm.f.Handshake(relay)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Check the relay HostInfo to see if we already established a relay through
|
||||||
|
existingRelay, ok := relayHostInfo.relayState.QueryRelayForByIp(vpnIp)
|
||||||
|
if !ok {
|
||||||
|
// No relays exist or requested yet.
|
||||||
|
if relayHostInfo.remote.IsValid() {
|
||||||
|
idx, err := AddRelay(hm.l, relayHostInfo, hm.mainHostMap, vpnIp, nil, TerminalType, Requested)
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(hm.l).Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
m := NebulaControl{
|
||||||
|
Type: NebulaControl_CreateRelayRequest,
|
||||||
|
InitiatorRelayIndex: idx,
|
||||||
|
}
|
||||||
|
|
||||||
|
switch relayHostInfo.GetCert().Certificate.Version() {
|
||||||
|
case cert.Version1:
|
||||||
|
if !hm.f.myVpnAddrs[0].Is4() {
|
||||||
|
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnIp.Is4() {
|
||||||
|
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
b := hm.f.myVpnAddrs[0].As4()
|
||||||
|
m.OldRelayFromAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
b = vpnIp.As4()
|
||||||
|
m.OldRelayToAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
case cert.Version2:
|
||||||
|
m.RelayFromAddr = netAddrToProtoAddr(hm.f.myVpnAddrs[0])
|
||||||
|
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
||||||
|
default:
|
||||||
|
hostinfo.logger(hm.l).Error("Unknown certificate version found while creating relay")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
msg, err := m.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(hm.l).Error("Failed to marshal Control message to create relay", "error", err)
|
||||||
|
} else {
|
||||||
|
hm.f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||||
|
hm.l.Info("send CreateRelayRequest",
|
||||||
|
"relayFrom", hm.f.myVpnAddrs[0],
|
||||||
|
"relayTo", vpnIp,
|
||||||
|
"initiatorRelayIndex", idx,
|
||||||
|
"relay", relay,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
switch existingRelay.State {
|
||||||
|
case Established:
|
||||||
|
hostinfo.logger(hm.l).Info("Send handshake via relay", "relay", relay.String())
|
||||||
|
hm.f.SendVia(relayHostInfo, existingRelay, hostinfo.HandshakePacket[0], make([]byte, 12), make([]byte, mtu), false)
|
||||||
|
case Disestablished:
|
||||||
|
// Mark this relay as 'requested'
|
||||||
|
relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested)
|
||||||
|
fallthrough
|
||||||
|
case Requested:
|
||||||
|
hostinfo.logger(hm.l).Info("Re-send CreateRelay request", "relay", relay.String())
|
||||||
|
// Re-send the CreateRelay request, in case the previous one was lost.
|
||||||
|
m := NebulaControl{
|
||||||
|
Type: NebulaControl_CreateRelayRequest,
|
||||||
|
InitiatorRelayIndex: existingRelay.LocalIndex,
|
||||||
|
}
|
||||||
|
|
||||||
|
switch relayHostInfo.GetCert().Certificate.Version() {
|
||||||
|
case cert.Version1:
|
||||||
|
if !hm.f.myVpnAddrs[0].Is4() {
|
||||||
|
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnIp.Is4() {
|
||||||
|
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
b := hm.f.myVpnAddrs[0].As4()
|
||||||
|
m.OldRelayFromAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
b = vpnIp.As4()
|
||||||
|
m.OldRelayToAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
case cert.Version2:
|
||||||
|
m.RelayFromAddr = netAddrToProtoAddr(hm.f.myVpnAddrs[0])
|
||||||
|
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
||||||
|
default:
|
||||||
|
hostinfo.logger(hm.l).Error("Unknown certificate version found while creating relay")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
msg, err := m.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(hm.l).Error("Failed to marshal Control message to create relay", "error", err)
|
||||||
|
} else {
|
||||||
|
// This must send over the hostinfo, not over hm.Hosts[ip]
|
||||||
|
hm.f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||||
|
hm.l.Info("send CreateRelayRequest",
|
||||||
|
"relayFrom", hm.f.myVpnAddrs[0],
|
||||||
|
"relayTo", vpnIp,
|
||||||
|
"initiatorRelayIndex", existingRelay.LocalIndex,
|
||||||
|
"relay", relay,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
case PeerRequested:
|
||||||
|
// PeerRequested only occurs in Forwarding relays, not Terminal relays, and this is a Terminal relay case.
|
||||||
|
fallthrough
|
||||||
|
default:
|
||||||
|
hostinfo.logger(hm.l).Error("Relay unexpected state",
|
||||||
|
"vpnIp", vpnIp,
|
||||||
|
"state", existingRelay.State,
|
||||||
|
"relay", relay,
|
||||||
|
)
|
||||||
|
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add
|
// If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add
|
||||||
if !lighthouseTriggered {
|
if !lighthouseTriggered {
|
||||||
@@ -436,11 +516,14 @@ func (hm *HandshakeManager) CheckAndComplete(hostinfo *HostInfo, handshakePacket
|
|||||||
// Check if we already have a tunnel with this vpn ip
|
// Check if we already have a tunnel with this vpn ip
|
||||||
existingHostInfo, found := hm.mainHostMap.Hosts[hostinfo.vpnAddrs[0]]
|
existingHostInfo, found := hm.mainHostMap.Hosts[hostinfo.vpnAddrs[0]]
|
||||||
if found && existingHostInfo != nil {
|
if found && existingHostInfo != nil {
|
||||||
// Is it just a delayed handshake packet? Check every hostinfo we hold for this address.
|
testHostInfo := existingHostInfo
|
||||||
for _, testHostInfo := range hm.mainHostMap.unlockedGetHostList(hostinfo.vpnAddrs[0]) {
|
for testHostInfo != nil {
|
||||||
|
// Is it just a delayed handshake packet?
|
||||||
if bytes.Equal(hostinfo.HandshakePacket[handshakePacket], testHostInfo.HandshakePacket[handshakePacket]) {
|
if bytes.Equal(hostinfo.HandshakePacket[handshakePacket], testHostInfo.HandshakePacket[handshakePacket]) {
|
||||||
return testHostInfo, ErrAlreadySeen
|
return testHostInfo, ErrAlreadySeen
|
||||||
}
|
}
|
||||||
|
|
||||||
|
testHostInfo = testHostInfo.next
|
||||||
}
|
}
|
||||||
|
|
||||||
// Is this a newer handshake?
|
// Is this a newer handshake?
|
||||||
@@ -468,6 +551,7 @@ func (hm *HandshakeManager) CheckAndComplete(hostinfo *HostInfo, handshakePacket
|
|||||||
// We have a collision, but this can happen since we can't control
|
// We have a collision, but this can happen since we can't control
|
||||||
// the remote ID. Just log about the situation as a note.
|
// the remote ID. Just log about the situation as a note.
|
||||||
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
||||||
|
"remoteIndex", hostinfo.remoteIndexId,
|
||||||
"collision", existingRemoteIndex.vpnAddrs,
|
"collision", existingRemoteIndex.vpnAddrs,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -490,6 +574,7 @@ func (hm *HandshakeManager) Complete(hostinfo *HostInfo, f *Interface) {
|
|||||||
// We have a collision, but this can happen since we can't control
|
// We have a collision, but this can happen since we can't control
|
||||||
// the remote ID. Just log about the situation as a note.
|
// the remote ID. Just log about the situation as a note.
|
||||||
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
||||||
|
"remoteIndex", hostinfo.remoteIndexId,
|
||||||
"collision", existingRemoteIndex.vpnAddrs,
|
"collision", existingRemoteIndex.vpnAddrs,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -502,7 +587,7 @@ func (hm *HandshakeManager) Complete(hostinfo *HostInfo, f *Interface) {
|
|||||||
// allocateIndex generates a unique localIndexId for this HostInfo
|
// allocateIndex generates a unique localIndexId for this HostInfo
|
||||||
// and adds it to the pendingHostMap. Will error if we are unable to generate
|
// and adds it to the pendingHostMap. Will error if we are unable to generate
|
||||||
// a unique localIndexId
|
// a unique localIndexId
|
||||||
func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) (uint32, error) {
|
func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) error {
|
||||||
hm.mainHostMap.RLock()
|
hm.mainHostMap.RLock()
|
||||||
defer hm.mainHostMap.RUnlock()
|
defer hm.mainHostMap.RUnlock()
|
||||||
hm.Lock()
|
hm.Lock()
|
||||||
@@ -511,7 +596,7 @@ func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) (uint32, error)
|
|||||||
for range 32 {
|
for range 32 {
|
||||||
index, err := generateIndex(hm.l)
|
index, err := generateIndex(hm.l)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
_, inPending := hm.indexes[index]
|
_, inPending := hm.indexes[index]
|
||||||
@@ -520,11 +605,11 @@ func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) (uint32, error)
|
|||||||
if !inMain && !inPending {
|
if !inMain && !inPending {
|
||||||
hh.hostinfo.localIndexId = index
|
hh.hostinfo.localIndexId = index
|
||||||
hm.indexes[index] = hh
|
hm.indexes[index] = hh
|
||||||
return index, nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return 0, errors.New("failed to generate unique localIndexId")
|
return errors.New("failed to generate unique localIndexId")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
||||||
@@ -535,10 +620,8 @@ func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
|||||||
|
|
||||||
func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
||||||
for _, addr := range hostinfo.vpnAddrs {
|
for _, addr := range hostinfo.vpnAddrs {
|
||||||
if cur, ok := hm.vpnIps[addr]; ok && cur.hostinfo == hostinfo {
|
|
||||||
delete(hm.vpnIps, addr)
|
delete(hm.vpnIps, addr)
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
if len(hm.vpnIps) == 0 {
|
if len(hm.vpnIps) == 0 {
|
||||||
hm.vpnIps = map[netip.Addr]*HandshakeHostInfo{}
|
hm.vpnIps = map[netip.Addr]*HandshakeHostInfo{}
|
||||||
@@ -645,522 +728,3 @@ func generateIndex(l *slog.Logger) (uint32, error) {
|
|||||||
func hsTimeout(tries int64, interval time.Duration) time.Duration {
|
func hsTimeout(tries int64, interval time.Duration) time.Duration {
|
||||||
return time.Duration(tries / 2 * ((2 * int64(interval)) + (tries-1)*int64(interval)))
|
return time.Duration(tries / 2 * ((2 * int64(interval)) + (tries-1)*int64(interval)))
|
||||||
}
|
}
|
||||||
|
|
||||||
// buildStage0Packet creates the initial handshake packet for the initiator.
|
|
||||||
func (hm *HandshakeManager) buildStage0Packet(hh *HandshakeHostInfo) bool {
|
|
||||||
cs := hm.f.pki.getCertState()
|
|
||||||
v := cs.DefaultVersion()
|
|
||||||
if hh.initiatingVersionOverride != cert.VersionPre1 {
|
|
||||||
v = hh.initiatingVersionOverride
|
|
||||||
} else if v < cert.Version2 {
|
|
||||||
for _, a := range hh.hostinfo.vpnAddrs {
|
|
||||||
if a.Is6() {
|
|
||||||
v = cert.Version2
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
cred := cs.GetCredential(v)
|
|
||||||
if cred == nil {
|
|
||||||
hm.f.l.Error("Unable to handshake with host because no certificate is available",
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs, "certVersion", v)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
machine, err := handshake.NewMachine(
|
|
||||||
v, cs.GetCredential,
|
|
||||||
hm.certVerifier(), func() (uint32, error) { return hm.allocateIndex(hh) },
|
|
||||||
true, header.HandshakeIXPSK0,
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
hm.f.l.Error("Failed to create handshake machine",
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs, "error", err)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
msg, err := machine.Initiate(nil)
|
|
||||||
if err != nil {
|
|
||||||
hm.f.l.Error("Failed to initiate handshake",
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs, "error", err)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
// hostinfo.ConnectionState stays nil until the handshake completes in
|
|
||||||
// continueHandshake. Pre-completion control surfaces guard with nil
|
|
||||||
// checks; the data plane never observes a pending hostinfo.
|
|
||||||
hh.hostinfo.HandshakePacket[handshakePacketStage0] = msg
|
|
||||||
hh.machine = machine
|
|
||||||
hh.ready = true
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// beginHandshake handles an incoming handshake packet that doesn't match any
|
|
||||||
// existing pending handshake. It creates a new responder Machine and processes
|
|
||||||
// the first message.
|
|
||||||
func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *header.H) {
|
|
||||||
f := hm.f
|
|
||||||
cs := f.pki.getCertState()
|
|
||||||
|
|
||||||
v := cs.DefaultVersion()
|
|
||||||
if cs.GetCredential(v) == nil {
|
|
||||||
f.l.Error("Unable to handshake with host because no certificate is available",
|
|
||||||
"from", via, "certVersion", v)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
machine, err := handshake.NewMachine(
|
|
||||||
v, cs.GetCredential,
|
|
||||||
hm.certVerifier(), func() (uint32, error) { return generateIndex(f.l) },
|
|
||||||
false, header.HandshakeIXPSK0,
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to create handshake machine", "from", via, "error", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
response, result, err := machine.ProcessPacket(nil, packet)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to process handshake packet", "from", via, "error", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if result == nil {
|
|
||||||
// Multi-message pattern: the responder Machine would need to be
|
|
||||||
// registered in hm.indexes so a future inbound packet finds it via
|
|
||||||
// continueHandshake. The current manager doesn't do that yet, so
|
|
||||||
// fail loudly rather than silently dropping the in-flight handshake.
|
|
||||||
// TODO: support multi-message responder flows (XX, pqIX, etc.).
|
|
||||||
// See also the IX-shaped cipher key assignment in handshake.Machine.
|
|
||||||
f.l.Error("multi-message handshake responder is not supported",
|
|
||||||
"from", via, "error", handshake.ErrMultiMessageUnsupported)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
remoteCert := result.RemoteCert
|
|
||||||
if remoteCert == nil {
|
|
||||||
f.l.Error("Handshake did not produce a peer certificate", "from", via)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Validate peer identity
|
|
||||||
vpnAddrs, anyVpnAddrsInCommon, ok := hm.validatePeerCert(via, remoteCert)
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo := &HostInfo{
|
|
||||||
ConnectionState: newConnectionStateFromResult(result),
|
|
||||||
localIndexId: result.LocalIndex,
|
|
||||||
remoteIndexId: result.RemoteIndex,
|
|
||||||
vpnAddrs: vpnAddrs,
|
|
||||||
HandshakePacket: make(map[uint8][]byte, 0),
|
|
||||||
lastHandshakeTime: result.HandshakeTime,
|
|
||||||
relayState: RelayState{
|
|
||||||
relays: nil,
|
|
||||||
relayForByAddr: map[netip.Addr]*Relay{},
|
|
||||||
relayForByIdx: map[uint32]*Relay{},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
msg := "Handshake message received"
|
|
||||||
if !anyVpnAddrsInCommon {
|
|
||||||
msg = "Handshake message received, but no vpnNetworks in common."
|
|
||||||
}
|
|
||||||
f.l.Info(msg,
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", remoteCert.Certificate.Name(),
|
|
||||||
"certVersion", remoteCert.Certificate.Version(),
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
"issuer", remoteCert.Certificate.Issuer(),
|
|
||||||
"initiatorIndex", result.RemoteIndex,
|
|
||||||
"responderIndex", result.LocalIndex,
|
|
||||||
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
|
||||||
)
|
|
||||||
|
|
||||||
// packet aliases the listener's incoming buffer, so this copy must stay.
|
|
||||||
hostinfo.HandshakePacket[handshakePacketStage0] = make([]byte, len(packet[header.Len:]))
|
|
||||||
copy(hostinfo.HandshakePacket[handshakePacketStage0], packet[header.Len:])
|
|
||||||
|
|
||||||
// response was freshly allocated by ProcessPacket; safe to retain directly.
|
|
||||||
if response != nil {
|
|
||||||
hostinfo.HandshakePacket[handshakePacketStage2] = response
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo.remotes = f.lightHouse.QueryCache(vpnAddrs)
|
|
||||||
if !via.IsRelayed {
|
|
||||||
hostinfo.SetRemote(via.UdpAddr)
|
|
||||||
}
|
|
||||||
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
|
||||||
|
|
||||||
existing, err := hm.CheckAndComplete(hostinfo, handshakePacketStage0, f)
|
|
||||||
if err != nil {
|
|
||||||
hm.handleCheckAndCompleteError(err, existing, hostinfo, via)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
hm.sendHandshakeResponse(via, response, hostinfo, false)
|
|
||||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
|
||||||
|
|
||||||
// Don't wait for UpdateWorker
|
|
||||||
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
|
||||||
f.lightHouse.TriggerUpdate()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// continueHandshake feeds an incoming packet to an existing pending handshake Machine.
|
|
||||||
func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostInfo, packet []byte) {
|
|
||||||
f := hm.f
|
|
||||||
|
|
||||||
hh.Lock()
|
|
||||||
defer hh.Unlock()
|
|
||||||
|
|
||||||
// Re-verify hh is still tracked. Between queryIndex returning and us taking
|
|
||||||
// hh.Lock, handleOutbound may have timed out and deleted it. Once we hold
|
|
||||||
// hh.Lock no other deleter can race our index: handleOutbound also takes
|
|
||||||
// hh.Lock first, and handleRecvError targets a main-hostmap entry with a
|
|
||||||
// different localIndexId.
|
|
||||||
hm.RLock()
|
|
||||||
cur, ok := hm.indexes[hh.hostinfo.localIndexId]
|
|
||||||
hm.RUnlock()
|
|
||||||
if !ok || cur != hh {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo := hh.hostinfo
|
|
||||||
if !via.IsRelayed {
|
|
||||||
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
|
|
||||||
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs, "from", via)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
machine := hh.machine
|
|
||||||
if machine == nil {
|
|
||||||
f.l.Error("No handshake machine available for continuation",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs, "from", via)
|
|
||||||
hm.DeleteHostInfo(hostinfo)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
response, result, err := machine.ProcessPacket(nil, packet)
|
|
||||||
if err != nil {
|
|
||||||
// Recoverable errors are routine noise, log at Debug. Fatal errors get a Warn.
|
|
||||||
if machine.Failed() {
|
|
||||||
f.l.Warn("Failed to process handshake packet, abandoning",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs, "from", via, "error", err)
|
|
||||||
hm.DeleteHostInfo(hostinfo)
|
|
||||||
} else {
|
|
||||||
f.l.Debug("Failed to process handshake packet",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs, "from", via, "error", err)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if response != nil {
|
|
||||||
hm.sendHandshakeResponse(via, response, hostinfo, false)
|
|
||||||
}
|
|
||||||
|
|
||||||
if result == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Handshake complete; build the ConnectionState now that we have keys and a verified peer cert.
|
|
||||||
hostinfo.ConnectionState = newConnectionStateFromResult(result)
|
|
||||||
|
|
||||||
remoteCert := result.RemoteCert
|
|
||||||
if remoteCert == nil {
|
|
||||||
f.l.Error("Handshake completed without peer certificate",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs, "from", via)
|
|
||||||
hm.DeleteHostInfo(hostinfo)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
vpnNetworks := remoteCert.Certificate.Networks()
|
|
||||||
hostinfo.remoteIndexId = result.RemoteIndex
|
|
||||||
hostinfo.lastHandshakeTime = result.HandshakeTime
|
|
||||||
|
|
||||||
if !via.IsRelayed {
|
|
||||||
hostinfo.SetRemote(via.UdpAddr)
|
|
||||||
} else {
|
|
||||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify correct host responded (initiator check)
|
|
||||||
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
|
||||||
correctHostResponded := false
|
|
||||||
anyVpnAddrsInCommon := false
|
|
||||||
for i, network := range vpnNetworks {
|
|
||||||
// inside.go drops self-routed packets at the firewall stage, but we'd
|
|
||||||
// rather not let a self-handshake complete in the first place: it
|
|
||||||
// wastes a hostmap slot, suppresses no log, and obscures routing
|
|
||||||
// misconfig. Explicit refusal here mirrors the responder-side check
|
|
||||||
// in validatePeerCert.
|
|
||||||
if f.myVpnAddrsTable.Contains(network.Addr()) {
|
|
||||||
f.l.Error("Refusing to handshake with myself",
|
|
||||||
"vpnNetworks", vpnNetworks,
|
|
||||||
"from", via,
|
|
||||||
"certName", remoteCert.Certificate.Name(),
|
|
||||||
"certVersion", remoteCert.Certificate.Version(),
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
"issuer", remoteCert.Certificate.Issuer(),
|
|
||||||
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
|
||||||
)
|
|
||||||
hm.DeleteHostInfo(hostinfo)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
vpnAddrs[i] = network.Addr()
|
|
||||||
if hostinfo.vpnAddrs[0] == network.Addr() {
|
|
||||||
correctHostResponded = true
|
|
||||||
}
|
|
||||||
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
|
||||||
anyVpnAddrsInCommon = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !correctHostResponded {
|
|
||||||
f.l.Info("Incorrect host responded to handshake",
|
|
||||||
"intendedVpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"haveVpnNetworks", vpnNetworks,
|
|
||||||
"from", via,
|
|
||||||
"certName", remoteCert.Certificate.Name(),
|
|
||||||
"certVersion", remoteCert.Certificate.Version(),
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
"issuer", remoteCert.Certificate.Issuer(),
|
|
||||||
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
|
||||||
)
|
|
||||||
|
|
||||||
hm.DeleteHostInfo(hostinfo)
|
|
||||||
hm.StartHandshake(hostinfo.vpnAddrs[0], func(newHH *HandshakeHostInfo) {
|
|
||||||
newHH.hostinfo.remotes = hostinfo.remotes
|
|
||||||
newHH.hostinfo.remotes.BlockRemote(via)
|
|
||||||
newHH.packetStore = hh.packetStore
|
|
||||||
hh.packetStore = []*cachedPacket{}
|
|
||||||
hostinfo.vpnAddrs = vpnAddrs
|
|
||||||
f.sendCloseTunnel(hostinfo)
|
|
||||||
})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
duration := time.Since(hh.startTime).Nanoseconds()
|
|
||||||
msg := "Handshake message received"
|
|
||||||
if !anyVpnAddrsInCommon {
|
|
||||||
msg = "Handshake message received, but no vpnNetworks in common."
|
|
||||||
}
|
|
||||||
f.l.Info(msg,
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", remoteCert.Certificate.Name(),
|
|
||||||
"certVersion", remoteCert.Certificate.Version(),
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
"issuer", remoteCert.Certificate.Issuer(),
|
|
||||||
"initiatorIndex", result.LocalIndex,
|
|
||||||
"responderIndex", result.RemoteIndex,
|
|
||||||
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
|
||||||
"durationNs", duration,
|
|
||||||
"sentCachedPackets", len(hh.packetStore),
|
|
||||||
)
|
|
||||||
|
|
||||||
hostinfo.vpnAddrs = vpnAddrs
|
|
||||||
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
|
||||||
|
|
||||||
hm.Complete(hostinfo, f)
|
|
||||||
|
|
||||||
if len(hh.packetStore) > 0 {
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hostinfo.logger(f.l).Debug("Sending stored packets", "count", len(hh.packetStore))
|
|
||||||
}
|
|
||||||
nb := make([]byte, 12, 12)
|
|
||||||
out := make([]byte, mtu)
|
|
||||||
for _, cp := range hh.packetStore {
|
|
||||||
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
|
||||||
}
|
|
||||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
|
||||||
f.metricHandshakes.Update(duration)
|
|
||||||
|
|
||||||
// Don't wait for UpdateWorker
|
|
||||||
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
|
||||||
f.lightHouse.TriggerUpdate()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// validatePeerCert checks the peer certificate for self-connection and remote allow list.
|
|
||||||
// Returns the VPN addrs, whether any of them fall within one of our own VPN
|
|
||||||
// networks, and true if valid; false if rejected.
|
|
||||||
func (hm *HandshakeManager) validatePeerCert(via ViaSender, remoteCert *cert.CachedCertificate) ([]netip.Addr, bool, bool) {
|
|
||||||
f := hm.f
|
|
||||||
vpnNetworks := remoteCert.Certificate.Networks()
|
|
||||||
|
|
||||||
// The cert package rejects host certs with no networks at parse time, so
|
|
||||||
// reaching this state would mean an invariant was bypassed elsewhere.
|
|
||||||
// Refuse explicitly so downstream code (which indexes vpnAddrs[0]) can't
|
|
||||||
// panic if that invariant ever changes.
|
|
||||||
if len(vpnNetworks) == 0 {
|
|
||||||
f.l.Info("No networks in certificate",
|
|
||||||
"from", via, "cert", remoteCert)
|
|
||||||
return nil, false, false
|
|
||||||
}
|
|
||||||
|
|
||||||
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
|
||||||
anyVpnAddrsInCommon := false
|
|
||||||
|
|
||||||
for i, network := range vpnNetworks {
|
|
||||||
if f.myVpnAddrsTable.Contains(network.Addr()) {
|
|
||||||
f.l.Error("Refusing to handshake with myself",
|
|
||||||
"vpnNetworks", vpnNetworks,
|
|
||||||
"from", via,
|
|
||||||
"certName", remoteCert.Certificate.Name(),
|
|
||||||
"certVersion", remoteCert.Certificate.Version(),
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
"issuer", remoteCert.Certificate.Issuer(),
|
|
||||||
)
|
|
||||||
return nil, false, false
|
|
||||||
}
|
|
||||||
vpnAddrs[i] = network.Addr()
|
|
||||||
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
|
||||||
anyVpnAddrsInCommon = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !via.IsRelayed {
|
|
||||||
if !f.lightHouse.GetRemoteAllowList().AllowAll(vpnAddrs, via.UdpAddr.Addr()) {
|
|
||||||
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
|
||||||
"vpnAddrs", vpnAddrs, "from", via)
|
|
||||||
return nil, false, false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return vpnAddrs, anyVpnAddrsInCommon, true
|
|
||||||
}
|
|
||||||
|
|
||||||
// sendHandshakeResponse sends a handshake response via the appropriate transport.
|
|
||||||
// cached is true when msg is a stored response being retransmitted because
|
|
||||||
// the peer's stage-1 retransmit landed (the ErrAlreadySeen path); false on a
|
|
||||||
// fresh response.
|
|
||||||
func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hostinfo *HostInfo, cached bool) {
|
|
||||||
if msg == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
f := hm.f
|
|
||||||
f.messageMetrics.Tx(header.Handshake, header.MessageSubType(msg[1]), 1)
|
|
||||||
|
|
||||||
// Common log fields. peerCert may be nil during intermediate
|
|
||||||
// multi-message flows (handshake hasn't completed yet); skip the cert
|
|
||||||
// block if so.
|
|
||||||
logFields := []any{
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"handshake", m{"stage": uint64(2), "style": header.SubTypeName(header.Handshake, header.HandshakeIXPSK0)},
|
|
||||||
"cached", cached,
|
|
||||||
"initiatorIndex", hostinfo.remoteIndexId,
|
|
||||||
"responderIndex", hostinfo.localIndexId,
|
|
||||||
}
|
|
||||||
if peerCert := hostinfo.ConnectionState.peerCert; peerCert != nil {
|
|
||||||
logFields = append(logFields,
|
|
||||||
"certName", peerCert.Certificate.Name(),
|
|
||||||
"certVersion", peerCert.Certificate.Version(),
|
|
||||||
"fingerprint", peerCert.Fingerprint,
|
|
||||||
"issuer", peerCert.Certificate.Issuer(),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !via.IsRelayed {
|
|
||||||
fields := append(logFields, "from", via)
|
|
||||||
err := f.outside.WriteTo(msg, via.UdpAddr)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to send handshake message", append(fields, "error", err)...)
|
|
||||||
} else {
|
|
||||||
f.l.Info("Handshake message sent", fields...)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
if via.relay == nil {
|
|
||||||
f.l.Error("Handshake send failed: both addr and via.relay are nil.")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
|
||||||
// We received a valid handshake on this relay, so make sure the relay
|
|
||||||
// state reflects that, in case it had been marked Disestablished.
|
|
||||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established)
|
|
||||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
|
||||||
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleCheckAndCompleteError handles errors from CheckAndComplete.
|
|
||||||
// This only fires from the responder-side beginHandshake path, after the
|
|
||||||
// peer cert has been validated and ConnectionState populated, so peerCert
|
|
||||||
// is always non-nil for the cases that log it.
|
|
||||||
func (hm *HandshakeManager) handleCheckAndCompleteError(err error, existing, hostinfo *HostInfo, via ViaSender) {
|
|
||||||
f := hm.f
|
|
||||||
peerCert := hostinfo.ConnectionState.peerCert
|
|
||||||
hsFields := m{"stage": uint64(1), "style": header.SubTypeName(header.Handshake, header.HandshakeIXPSK0)}
|
|
||||||
|
|
||||||
switch err {
|
|
||||||
case ErrAlreadySeen:
|
|
||||||
if existing.SetRemoteIfPreferred(f.hostMap, via) {
|
|
||||||
f.SendMessageToVpnAddr(header.Test, header.TestRequest, hostinfo.vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
|
||||||
}
|
|
||||||
// Resend the original response. The peer is committed to that response's
|
|
||||||
// ephemeral keys; a freshly-built one would have different keys and break
|
|
||||||
// the tunnel even though both sides "completed" the handshake.
|
|
||||||
if msg := existing.HandshakePacket[handshakePacketStage2]; msg != nil {
|
|
||||||
hm.sendHandshakeResponse(via, msg, existing, true)
|
|
||||||
}
|
|
||||||
|
|
||||||
case ErrExistingHostInfo:
|
|
||||||
f.l.Info("Handshake too old",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", peerCert.Certificate.Name(),
|
|
||||||
"certVersion", peerCert.Certificate.Version(),
|
|
||||||
"fingerprint", peerCert.Fingerprint,
|
|
||||||
"issuer", peerCert.Certificate.Issuer(),
|
|
||||||
"oldHandshakeTime", existing.lastHandshakeTime,
|
|
||||||
"newHandshakeTime", hostinfo.lastHandshakeTime,
|
|
||||||
"initiatorIndex", hostinfo.remoteIndexId,
|
|
||||||
"responderIndex", hostinfo.localIndexId,
|
|
||||||
"handshake", hsFields,
|
|
||||||
)
|
|
||||||
f.SendMessageToVpnAddr(header.Test, header.TestRequest, hostinfo.vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
|
||||||
|
|
||||||
case ErrLocalIndexCollision:
|
|
||||||
f.l.Error("Failed to add HostInfo due to localIndex collision",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", peerCert.Certificate.Name(),
|
|
||||||
"certVersion", peerCert.Certificate.Version(),
|
|
||||||
"fingerprint", peerCert.Fingerprint,
|
|
||||||
"issuer", peerCert.Certificate.Issuer(),
|
|
||||||
"localIndex", hostinfo.localIndexId,
|
|
||||||
"initiatorIndex", hostinfo.remoteIndexId,
|
|
||||||
"responderIndex", hostinfo.localIndexId,
|
|
||||||
"handshake", hsFields,
|
|
||||||
)
|
|
||||||
|
|
||||||
default:
|
|
||||||
f.l.Error("Failed to add HostInfo to HostMap",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"error", err,
|
|
||||||
"certName", peerCert.Certificate.Name(),
|
|
||||||
"certVersion", peerCert.Certificate.Version(),
|
|
||||||
"fingerprint", peerCert.Fingerprint,
|
|
||||||
"issuer", peerCert.Certificate.Issuer(),
|
|
||||||
"initiatorIndex", hostinfo.remoteIndexId,
|
|
||||||
"responderIndex", hostinfo.localIndexId,
|
|
||||||
"handshake", hsFields,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// certVerifier returns a CertVerifier that validates certs against the current CA pool.
|
|
||||||
func (hm *HandshakeManager) certVerifier() handshake.CertVerifier {
|
|
||||||
return func(c cert.Certificate) (*cert.CachedCertificate, error) {
|
|
||||||
return hm.f.pki.GetCAPool().VerifyCertificate(time.Now(), c)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
+1
-136
@@ -5,7 +5,6 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
@@ -28,7 +27,7 @@ func Test_NewHandshakeManagerVpnIp(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
blah := NewHandshakeManager(l, mainHM, lh, &udp.NoopConn{}, defaultHandshakeConfig)
|
blah := NewHandshakeManager(l, mainHM, lh, &udp.NoopConn{}, defaultHandshakeConfig)
|
||||||
@@ -101,137 +100,3 @@ func (mw *mockEncWriter) GetHostInfo(_ netip.Addr) *HostInfo {
|
|||||||
func (mw *mockEncWriter) GetCertState() *CertState {
|
func (mw *mockEncWriter) GetCertState() *CertState {
|
||||||
return &CertState{initiatingVersion: cert.Version2}
|
return &CertState{initiatingVersion: cert.Version2}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestValidatePeerCert(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
|
|
||||||
myNetwork := netip.MustParsePrefix("10.0.0.1/24")
|
|
||||||
myAddrTable := new(bart.Lite)
|
|
||||||
myAddrTable.Insert(netip.PrefixFrom(myNetwork.Addr(), myNetwork.Addr().BitLen()))
|
|
||||||
myNetTable := new(bart.Lite)
|
|
||||||
myNetTable.Insert(myNetwork.Masked())
|
|
||||||
|
|
||||||
newHM := func() *HandshakeManager {
|
|
||||||
hm := NewHandshakeManager(l, newHostMap(l), newTestLighthouse(), &udp.NoopConn{}, defaultHandshakeConfig)
|
|
||||||
hm.f = &Interface{
|
|
||||||
handshakeManager: hm,
|
|
||||||
pki: &PKI{},
|
|
||||||
l: l,
|
|
||||||
myVpnAddrsTable: myAddrTable,
|
|
||||||
myVpnNetworksTable: myNetTable,
|
|
||||||
lightHouse: hm.lightHouse,
|
|
||||||
}
|
|
||||||
return hm
|
|
||||||
}
|
|
||||||
|
|
||||||
cached := func(networks ...netip.Prefix) *cert.CachedCertificate {
|
|
||||||
return &cert.CachedCertificate{
|
|
||||||
Certificate: &dummyCert{name: "peer", networks: networks},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
via := ViaSender{
|
|
||||||
UdpAddr: netip.MustParseAddrPort("198.51.100.7:4242"),
|
|
||||||
IsRelayed: true, // skip the remote allow list (covered separately)
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Run("addr inside our networks sets anyVpnAddrsInCommon", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
// 10.0.0.2 falls inside our 10.0.0.0/24
|
|
||||||
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("10.0.0.2/24")))
|
|
||||||
assert.True(t, ok)
|
|
||||||
assert.True(t, common)
|
|
||||||
assert.Equal(t, []netip.Addr{netip.MustParseAddr("10.0.0.2")}, addrs)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("addr outside our networks leaves anyVpnAddrsInCommon false", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("192.168.1.5/24")))
|
|
||||||
assert.True(t, ok)
|
|
||||||
assert.False(t, common)
|
|
||||||
assert.Equal(t, []netip.Addr{netip.MustParseAddr("192.168.1.5")}, addrs)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("any matching network is enough", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
addrs, common, ok := hm.validatePeerCert(via, cached(
|
|
||||||
netip.MustParsePrefix("192.168.1.5/24"),
|
|
||||||
netip.MustParsePrefix("10.0.0.42/24"),
|
|
||||||
))
|
|
||||||
assert.True(t, ok)
|
|
||||||
assert.True(t, common)
|
|
||||||
assert.Len(t, addrs, 2)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("self-handshake is rejected", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
// 10.0.0.1 is in myVpnAddrsTable
|
|
||||||
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("10.0.0.1/24")))
|
|
||||||
assert.False(t, ok)
|
|
||||||
assert.False(t, common)
|
|
||||||
assert.Nil(t, addrs)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("cert with no networks is rejected", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
addrs, common, ok := hm.validatePeerCert(via, cached())
|
|
||||||
assert.False(t, ok)
|
|
||||||
assert.False(t, common)
|
|
||||||
assert.Nil(t, addrs)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandleIncomingDispatch(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
|
|
||||||
newHM := func() *HandshakeManager {
|
|
||||||
hm := NewHandshakeManager(l, newHostMap(l), newTestLighthouse(), &udp.NoopConn{}, defaultHandshakeConfig)
|
|
||||||
hm.f = &Interface{
|
|
||||||
handshakeManager: hm,
|
|
||||||
pki: &PKI{},
|
|
||||||
l: l,
|
|
||||||
}
|
|
||||||
return hm
|
|
||||||
}
|
|
||||||
|
|
||||||
via := ViaSender{
|
|
||||||
UdpAddr: netip.MustParseAddrPort("198.51.100.7:4242"),
|
|
||||||
IsRelayed: true, // bypass remote allow list
|
|
||||||
}
|
|
||||||
|
|
||||||
// A packet body of zero length is fine for these tests: dispatch is
|
|
||||||
// gated on header fields, and we assert that we never reach noise/cert
|
|
||||||
// processing for any of the malformed shapes here.
|
|
||||||
pkt := make([]byte, header.Len)
|
|
||||||
|
|
||||||
t.Run("unsupported subtype dropped", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
h := &header.H{Type: header.Handshake, Subtype: header.MessageSubType(99), MessageCounter: 1}
|
|
||||||
hm.HandleIncoming(via, pkt, h)
|
|
||||||
assert.Empty(t, hm.indexes, "no pending handshake should be created")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("stage-1 with non-zero RemoteIndex dropped", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
h := &header.H{
|
|
||||||
Type: header.Handshake,
|
|
||||||
Subtype: header.HandshakeIXPSK0,
|
|
||||||
RemoteIndex: 0xdeadbeef,
|
|
||||||
MessageCounter: 1,
|
|
||||||
}
|
|
||||||
hm.HandleIncoming(via, pkt, h)
|
|
||||||
assert.Empty(t, hm.indexes, "spoofed stage-1 must not create a pending machine")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("continuation with no matching pending index dropped", func(t *testing.T) {
|
|
||||||
hm := newHM()
|
|
||||||
h := &header.H{
|
|
||||||
Type: header.Handshake,
|
|
||||||
Subtype: header.HandshakeIXPSK0,
|
|
||||||
RemoteIndex: 0xcafef00d,
|
|
||||||
MessageCounter: 2,
|
|
||||||
}
|
|
||||||
hm.HandleIncoming(via, pkt, h)
|
|
||||||
assert.Empty(t, hm.indexes, "orphan stage-2 must not create state")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -174,10 +174,6 @@ func (h *H) SubTypeName() string {
|
|||||||
return SubTypeName(h.Type, h.Subtype)
|
return SubTypeName(h.Type, h.Subtype)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *H) IsValidSubType() bool {
|
|
||||||
return IsValidSubType(h.Type, h.Subtype)
|
|
||||||
}
|
|
||||||
|
|
||||||
// SubTypeName will transform a nebula message sub type into a human string
|
// SubTypeName will transform a nebula message sub type into a human string
|
||||||
func SubTypeName(t MessageType, s MessageSubType) string {
|
func SubTypeName(t MessageType, s MessageSubType) string {
|
||||||
if n, ok := subTypeMap[t]; ok {
|
if n, ok := subTypeMap[t]; ok {
|
||||||
@@ -189,16 +185,6 @@ func SubTypeName(t MessageType, s MessageSubType) string {
|
|||||||
return "unknown"
|
return "unknown"
|
||||||
}
|
}
|
||||||
|
|
||||||
func IsValidSubType(t MessageType, s MessageSubType) bool {
|
|
||||||
if n, ok := subTypeMap[t]; ok {
|
|
||||||
if _, ok := (*n)[s]; ok {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewHeader turns bytes into a header
|
// NewHeader turns bytes into a header
|
||||||
func NewHeader(b []byte) (*H, error) {
|
func NewHeader(b []byte) (*H, error) {
|
||||||
h := new(H)
|
h := new(H)
|
||||||
|
|||||||
+102
-154
@@ -60,16 +60,7 @@ type HostMap struct {
|
|||||||
Indexes map[uint32]*HostInfo
|
Indexes map[uint32]*HostInfo
|
||||||
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
|
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
|
||||||
RemoteIndexes map[uint32]*HostInfo
|
RemoteIndexes map[uint32]*HostInfo
|
||||||
// Hosts maps a vpn address to its primary hostinfo, one entry per address we hold a tunnel
|
|
||||||
// for. moreHosts only has an entry while an address is held by 2 or more hostinfos and stores
|
|
||||||
// the full most-recent-first list; moreHosts[a][0] is always the same hostinfo as Hosts[a].
|
|
||||||
// Each address gets its own independent list, so a hostinfo owning multiple addresses can
|
|
||||||
// never corrupt another address's ordering the way the old shared next/prev chain could.
|
|
||||||
// Entries in moreHosts are only ever written by unlockedSetHostsForAddr; Hosts is written
|
|
||||||
// directly only in the single-hostinfo fast paths where moreHosts is known to have no entry,
|
|
||||||
// and unlockedDeleteHostInfo swaps either map for a fresh one when it fully drains.
|
|
||||||
Hosts map[netip.Addr]*HostInfo
|
Hosts map[netip.Addr]*HostInfo
|
||||||
moreHosts map[netip.Addr][]*HostInfo
|
|
||||||
preferredRanges atomic.Pointer[[]netip.Prefix]
|
preferredRanges atomic.Pointer[[]netip.Prefix]
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
@@ -147,9 +138,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 +229,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 +266,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 +282,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
|
||||||
}
|
}
|
||||||
@@ -338,7 +334,6 @@ func newHostMap(l *slog.Logger) *HostMap {
|
|||||||
Relays: map[uint32]*HostInfo{},
|
Relays: map[uint32]*HostInfo{},
|
||||||
RemoteIndexes: map[uint32]*HostInfo{},
|
RemoteIndexes: map[uint32]*HostInfo{},
|
||||||
Hosts: map[netip.Addr]*HostInfo{},
|
Hosts: map[netip.Addr]*HostInfo{},
|
||||||
moreHosts: map[netip.Addr][]*HostInfo{},
|
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -387,55 +382,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,67 +400,86 @@ 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...)
|
// Unlink this hostinfo
|
||||||
hm.unlockedSetHostsForAddr(addr, list)
|
if hostinfo.prev != nil {
|
||||||
|
hostinfo.prev.next = hostinfo.next
|
||||||
}
|
}
|
||||||
return true
|
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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Go maps never shrink their buckets, replace fully drained maps so a node that churned
|
func (hm *HostMap) unlockedInnerDeleteHostInfo(hostinfo *HostInfo, addr netip.Addr) {
|
||||||
// through a large peer count gives the memory back. Same idiom as the index maps below.
|
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 {
|
if len(hm.Hosts) == 0 {
|
||||||
hm.Hosts = map[netip.Addr]*HostInfo{}
|
hm.Hosts = map[netip.Addr]*HostInfo{}
|
||||||
}
|
}
|
||||||
if len(hm.moreHosts) == 0 {
|
|
||||||
hm.moreHosts = 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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo.next = nil
|
||||||
|
hostinfo.prev = nil
|
||||||
|
|
||||||
// The remote index uses index ids outside our control so lets make sure we are only removing
|
// The remote index uses index ids outside our control so lets make sure we are only removing
|
||||||
// the remote index pointer here if it points to the hostinfo we are deleting
|
// the remote index pointer here if it points to the hostinfo we are deleting
|
||||||
hostinfo2, ok := hm.RemoteIndexes[hostinfo.remoteIndexId]
|
hostinfo2, ok := hm.RemoteIndexes[hostinfo.remoteIndexId]
|
||||||
@@ -530,7 +502,7 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
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 +511,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 +554,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 h != nil {
|
||||||
for _, targetIp := range targetIps {
|
for _, targetIp := range targetIps {
|
||||||
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
||||||
if ok && r.State == Established {
|
if ok && r.State == Established {
|
||||||
return h, r, nil
|
return h, r, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
h = h.next
|
||||||
if list, ok := hm.moreHosts[relayHostIp]; ok {
|
|
||||||
// list[0] is the primary we already checked
|
|
||||||
for _, h := range list[1:] {
|
|
||||||
for _, targetIp := range targetIps {
|
|
||||||
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
|
||||||
if ok && r.State == Established {
|
|
||||||
return h, r, nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil, nil, errors.New("unable to find host with relay")
|
return nil, nil, errors.New("unable to find host with relay")
|
||||||
@@ -615,14 +574,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 {
|
||||||
|
for h != nil {
|
||||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
||||||
|
h = h.next
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for _, rs := range hi.relayState.CopyAllRelayFor() {
|
for _, rs := range hi.relayState.CopyAllRelayFor() {
|
||||||
if rs.Type == ForwardingType {
|
if rs.Type == ForwardingType {
|
||||||
for _, h := range hm.unlockedGetHostList(rs.PeerAddr) {
|
if h, ok := hm.Hosts[rs.PeerAddr]; ok {
|
||||||
|
for h != nil {
|
||||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
||||||
|
h = h.next
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -658,11 +623,6 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
|
|||||||
hm.Indexes[hostinfo.localIndexId] = hostinfo
|
hm.Indexes[hostinfo.localIndexId] = hostinfo
|
||||||
hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo
|
hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo
|
||||||
|
|
||||||
hostinfo.out.Store(true)
|
|
||||||
if f.connectionManager != nil { // f.connectionManager is only nil in some unit tests
|
|
||||||
f.connectionManager.trafficTimer.Add(hostinfo.localIndexId, f.connectionManager.checkInterval)
|
|
||||||
}
|
|
||||||
|
|
||||||
if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hm.l.Debug("Hostmap vpnIp added",
|
hm.l.Debug("Hostmap vpnIp added",
|
||||||
"hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts),
|
"hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts),
|
||||||
@@ -672,27 +632,22 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) {
|
func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) {
|
||||||
existing, ok := hm.Hosts[vpnAddr]
|
existing := hm.Hosts[vpnAddr]
|
||||||
if !ok {
|
|
||||||
// Common case, the first hostinfo for this address. moreHosts stays empty.
|
|
||||||
hm.Hosts[vpnAddr] = hostinfo
|
hm.Hosts[vpnAddr] = hostinfo
|
||||||
return
|
|
||||||
|
if existing != nil && existing != hostinfo {
|
||||||
|
hostinfo.next = existing
|
||||||
|
existing.prev = hostinfo
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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 +679,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 +721,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 +737,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
|
||||||
|
|||||||
+138
-295
@@ -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,243 +104,94 @@ 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) {
|
||||||
|
|||||||
@@ -9,10 +9,11 @@ import (
|
|||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
err := newPacket(packet, false, fwPacket)
|
err := newPacket(packet, false, fwPacket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
@@ -57,7 +58,7 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
})
|
})
|
||||||
|
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
||||||
"vpnAddr", fwPacket.RemoteAddr,
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
@@ -73,10 +74,9 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
|
|
||||||
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
if dropReason == nil {
|
if dropReason == nil {
|
||||||
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
f.sendInsideMessage(hostinfo, packet, nb, sendBatch, rejectBuf, q)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
||||||
"fwPacket", fwPacket,
|
"fwPacket", fwPacket,
|
||||||
@@ -86,8 +86,69 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// sendInsideMessage encrypts a firewall-approved inside packet into the
|
||||||
|
// caller's batch slot for later sendmmsg flush. When hostinfo.remote is not
|
||||||
|
// valid we fall through to the relay slow path via the unbatched sendNoMetrics
|
||||||
|
// so relay behavior is unchanged.
|
||||||
|
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, p, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int) {
|
||||||
|
ci := hostinfo.ConnectionState
|
||||||
|
if ci.eKey == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if !hostinfo.remote.IsValid() {
|
||||||
|
// Slow path: relay fallback. Reuse rejectBuf as the ciphertext
|
||||||
|
// scratch; sendNoMetrics arranges header space for SendVia.
|
||||||
|
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, p, nb, rejectBuf, q)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
scratch := sendBatch.Next()
|
||||||
|
if scratch == nil {
|
||||||
|
// Batch full: bypass batching and send this packet directly so we
|
||||||
|
// never drop traffic on over-subscribed iterations.
|
||||||
|
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, p, nb, rejectBuf, q)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Lock()
|
||||||
|
}
|
||||||
|
c := ci.messageCounter.Add(1)
|
||||||
|
|
||||||
|
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
||||||
|
f.connectionManager.Out(hostinfo)
|
||||||
|
|
||||||
|
if hostinfo.lastRebindCount != f.rebindCount {
|
||||||
|
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
||||||
|
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||||
|
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
|
||||||
|
hostinfo.lastRebindCount = f.rebindCount
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind counter",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
out, err := ci.eKey.EncryptDanger(out, out, p, c, nb)
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Unlock()
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
|
||||||
|
"error", err,
|
||||||
|
"udpAddr", hostinfo.remote,
|
||||||
|
"counter", c,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
sendBatch.Commit(len(out), hostinfo.remote)
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||||
if !f.firewall.OutboundSendReject {
|
if !f.firewall.InSendReject {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -103,7 +164,7 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
|
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
|
||||||
if !f.firewall.InboundSendReject {
|
if !f.firewall.OutSendReject {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -333,7 +394,7 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
err = f.writers[0].WriteTo(out, via.GetRemote())
|
err = f.writers[0].WriteTo(out, via.remote)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
||||||
}
|
}
|
||||||
@@ -344,7 +405,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
|||||||
if ci.eKey == nil {
|
if ci.eKey == nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
useRelay := !remote.IsValid() && !hostinfo.GetRemote().IsValid()
|
useRelay := !remote.IsValid() && !hostinfo.remote.IsValid()
|
||||||
fullOut := out
|
fullOut := out
|
||||||
|
|
||||||
if useRelay {
|
if useRelay {
|
||||||
@@ -391,6 +452,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
|||||||
"error", err,
|
"error", err,
|
||||||
"udpAddr", remote,
|
"udpAddr", remote,
|
||||||
"counter", c,
|
"counter", c,
|
||||||
|
"attemptedCounter", c,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -403,12 +465,12 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
|||||||
"udpAddr", remote,
|
"udpAddr", remote,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
} else if hr := hostinfo.GetRemote(); hr.IsValid() {
|
} else if hostinfo.remote.IsValid() {
|
||||||
err = f.writers[q].WriteTo(out, hr)
|
err = f.writers[q].WriteTo(out, hostinfo.remote)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
||||||
"error", err,
|
"error", err,
|
||||||
"udpAddr", hr,
|
"udpAddr", remote,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
+67
-67
@@ -4,10 +4,8 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"slices"
|
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
@@ -15,11 +13,12 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -90,7 +89,11 @@ type Interface struct {
|
|||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
writers []udp.Conn
|
writers []udp.Conn
|
||||||
readers []io.ReadWriteCloser
|
readers []tio.Queue
|
||||||
|
// batchers is one per tun queue, wrapping readers[i].
|
||||||
|
// decryptToTun sends plaintext into the batch.RxBatcher;
|
||||||
|
// listenOut calls its Flush at the end of each UDP recvmmsg batch.
|
||||||
|
batchers []batch.RxBatcher
|
||||||
wg sync.WaitGroup
|
wg sync.WaitGroup
|
||||||
|
|
||||||
// fatalErr holds the first unexpected reader error that caused shutdown.
|
// fatalErr holds the first unexpected reader error that caused shutdown.
|
||||||
@@ -189,7 +192,8 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
routines: c.routines,
|
routines: c.routines,
|
||||||
version: c.version,
|
version: c.version,
|
||||||
writers: make([]udp.Conn, c.routines),
|
writers: make([]udp.Conn, c.routines),
|
||||||
readers: make([]io.ReadWriteCloser, c.routines),
|
readers: make([]tio.Queue, c.routines),
|
||||||
|
batchers: make([]batch.RxBatcher, c.routines),
|
||||||
myVpnNetworks: cs.myVpnNetworks,
|
myVpnNetworks: cs.myVpnNetworks,
|
||||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||||
myVpnAddrs: cs.myVpnAddrs,
|
myVpnAddrs: cs.myVpnAddrs,
|
||||||
@@ -215,9 +219,6 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
|
|
||||||
ifce.connectionManager.intf = ifce
|
ifce.connectionManager.intf = ifce
|
||||||
|
|
||||||
// Held until Close so waiting on the interface blocks until the resources are actually released
|
|
||||||
ifce.wg.Add(1)
|
|
||||||
|
|
||||||
return ifce, nil
|
return ifce, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -250,27 +251,29 @@ func (f *Interface) activate() error {
|
|||||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||||
|
|
||||||
// Prepare n tun queues
|
// Prepare n tun queues
|
||||||
var reader io.ReadWriteCloser = f.inside
|
|
||||||
for i := 0; i < f.routines; i++ {
|
for i := 0; i < f.routines; i++ {
|
||||||
if i > 0 {
|
if i > 0 {
|
||||||
reader, err = f.inside.NewMultiQueueReader()
|
if err = f.inside.NewMultiQueueReader(); err != nil {
|
||||||
if err != nil {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
f.readers[i] = reader
|
}
|
||||||
|
f.readers = f.inside.Readers()
|
||||||
|
for i := range f.readers {
|
||||||
|
f.batchers[i] = batch.NewTCPCoalescer(f.readers[i])
|
||||||
}
|
}
|
||||||
|
|
||||||
// On error the caller owns the cleanup, Control.Start cancels the service context
|
f.wg.Add(1) // for us to wait on Close() to return
|
||||||
// before releasing our resources so a waiter never observes a live context
|
|
||||||
if err = f.inside.Activate(); err != nil {
|
if err = f.inside.Activate(); err != nil {
|
||||||
|
f.wg.Done()
|
||||||
|
f.inside.Close()
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) run() {
|
func (f *Interface) run() (func() error, error) {
|
||||||
// Launch n queues to read packets from udp
|
// Launch n queues to read packets from udp
|
||||||
for i := 0; i < f.routines; i++ {
|
for i := 0; i < f.routines; i++ {
|
||||||
f.wg.Go(func() {
|
f.wg.Go(func() {
|
||||||
@@ -285,14 +288,13 @@ func (f *Interface) run() {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
return func() error {
|
||||||
|
|
||||||
func (f *Interface) wait() error {
|
|
||||||
f.wg.Wait()
|
f.wg.Wait()
|
||||||
if e := f.fatalErr.Load(); e != nil {
|
if e := f.fatalErr.Load(); e != nil {
|
||||||
return *e
|
return *e
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one
|
// onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one
|
||||||
@@ -316,19 +318,26 @@ func (f *Interface) listenOut(i int) {
|
|||||||
|
|
||||||
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
lhh := f.lightHouse.NewRequestHandler()
|
lhh := f.lightHouse.NewRequestHandler()
|
||||||
plaintext := make([]byte, udp.MTU)
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
coalescer := f.batchers[i]
|
||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
|
||||||
})
|
|
||||||
|
|
||||||
// An error after teardown began is shutdown noise, the closed flag covers resources
|
listener := func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||||
// Close releases itself and the cancelled ctx covers ones torn down by their owners
|
plaintext := f.batchers[i].Reserve(len(payload))
|
||||||
// reacting to it, like the user device pipes
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
||||||
if err != nil && !f.closed.Load() && f.ctx.Err() == nil {
|
}
|
||||||
|
|
||||||
|
flusher := func() {
|
||||||
|
if err := coalescer.Flush(); err != nil {
|
||||||
|
f.l.Error("Failed to flush tun coalescer", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err := li.ListenOut(listener, flusher)
|
||||||
|
|
||||||
|
if err != nil && !f.closed.Load() {
|
||||||
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
||||||
f.onFatal(err)
|
f.onFatal(err)
|
||||||
}
|
}
|
||||||
@@ -336,31 +345,47 @@ func (f *Interface) listenOut(i int) {
|
|||||||
f.l.Debug("underlay reader is done", "reader", i)
|
f.l.Debug("underlay reader is done", "reader", i)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
func (f *Interface) listenIn(reader tio.Queue, i int) {
|
||||||
packet := make([]byte, mtu)
|
rejectBuf := make([]byte, mtu)
|
||||||
out := make([]byte, mtu)
|
sb := batch.NewSendBatch(batch.SendBatchCap, udp.MTU+32)
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
n, err := reader.Read(packet)
|
pkts, err := reader.Read()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Same shutdown noise handling as listenOut
|
if !f.closed.Load() {
|
||||||
if !f.closed.Load() && f.ctx.Err() == nil {
|
|
||||||
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
||||||
f.onFatal(err)
|
f.onFatal(err)
|
||||||
}
|
}
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get())
|
sb.Reset()
|
||||||
|
for _, pkt := range pkts {
|
||||||
|
if sb.Len() >= sb.Cap() {
|
||||||
|
f.flushBatch(sb, i)
|
||||||
|
sb.Reset()
|
||||||
|
}
|
||||||
|
f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get())
|
||||||
|
}
|
||||||
|
if sb.Len() > 0 {
|
||||||
|
f.flushBatch(sb, i)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
f.l.Debug("overlay reader is done", "reader", i)
|
f.l.Debug("overlay reader is done", "reader", i)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (f *Interface) flushBatch(sb batch.TxBatcher, q int) {
|
||||||
|
bufs, dsts := sb.Get()
|
||||||
|
if err := f.writers[q].WriteBatch(bufs, dsts); err != nil {
|
||||||
|
f.l.Error("Failed to write outgoing batch", "error", err, "writer", q)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
||||||
c.RegisterReloadCallback(f.reloadFirewall)
|
c.RegisterReloadCallback(f.reloadFirewall)
|
||||||
c.RegisterReloadCallback(f.reloadSendRecvError)
|
c.RegisterReloadCallback(f.reloadSendRecvError)
|
||||||
@@ -384,22 +409,13 @@ func (f *Interface) reloadDisconnectInvalid(c *config.C) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) reloadFirewall(c *config.C) {
|
func (f *Interface) reloadFirewall(c *config.C) {
|
||||||
cs := f.pki.getCertState()
|
//TODO: need to trigger/detect if the certificate changed too
|
||||||
curCert := cs.getCertificate(cert.Version2)
|
if c.HasChanged("firewall") == false {
|
||||||
if curCert == nil {
|
|
||||||
curCert = cs.getCertificate(cert.Version1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The firewall builds its routableNetworks set from the certificate's UnsafeNetworks at construction.
|
|
||||||
// Check to see if that set has changed, and if so, rebuild the firewall.
|
|
||||||
certUnsafeChanged := curCert != nil && !slices.Equal(curCert.UnsafeNetworks(), f.firewall.unsafeNetworks)
|
|
||||||
|
|
||||||
if !c.HasChanged("firewall") && !certUnsafeChanged {
|
|
||||||
f.l.Debug("No firewall config change detected")
|
f.l.Debug("No firewall config change detected")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
fw, err := NewFirewallFromConfig(f.l, cs, c)
|
fw, err := NewFirewallFromConfig(f.l, f.pki.getCertState(), c)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Error while creating firewall during reload", "error", err)
|
f.l.Error("Error while creating firewall during reload", "error", err)
|
||||||
return
|
return
|
||||||
@@ -509,7 +525,11 @@ func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
|||||||
certInitiatingVersion := metrics.GetOrRegisterGauge("certificate.initiating_version", nil)
|
certInitiatingVersion := metrics.GetOrRegisterGauge("certificate.initiating_version", nil)
|
||||||
certMaxVersion := metrics.GetOrRegisterGauge("certificate.max_version", nil)
|
certMaxVersion := metrics.GetOrRegisterGauge("certificate.max_version", nil)
|
||||||
|
|
||||||
emit := func() {
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
f.firewall.EmitStats()
|
f.firewall.EmitStats()
|
||||||
f.handshakeManager.EmitStats()
|
f.handshakeManager.EmitStats()
|
||||||
udpStats()
|
udpStats()
|
||||||
@@ -526,18 +546,6 @@ func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
|||||||
certMaxVersion.Update(int64(certState.v1Cert.Version()))
|
certMaxVersion.Update(int64(certState.v1Cert.Version()))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Prime gauges so a Prometheus scrape that lands before the first tick
|
|
||||||
// sees real values instead of the zero defaults (issue #907).
|
|
||||||
emit()
|
|
||||||
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
return
|
|
||||||
case <-ticker.C:
|
|
||||||
emit()
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -549,15 +557,9 @@ func (f *Interface) GetCertState() *CertState {
|
|||||||
return f.pki.getCertState()
|
return f.pki.getCertState()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close releases the interface's resources: the udp sockets and the tun device.
|
|
||||||
// It is idempotent and safe to call at any point in the lifecycle, including on an interface that never activated,
|
|
||||||
// calls after the first return nil without doing anything.
|
|
||||||
func (f *Interface) Close() error {
|
func (f *Interface) Close() error {
|
||||||
if !f.closed.CompareAndSwap(false, true) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var errs []error
|
var errs []error
|
||||||
|
f.closed.Store(true)
|
||||||
|
|
||||||
// Release the udp readers
|
// Release the udp readers
|
||||||
for i, u := range f.writers {
|
for i, u := range f.writers {
|
||||||
@@ -573,8 +575,6 @@ func (f *Interface) Close() error {
|
|||||||
if closeErr != nil {
|
if closeErr != nil {
|
||||||
errs = append(errs, closeErr)
|
errs = append(errs, closeErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Release the construction token so waiters know the resources are gone
|
|
||||||
f.wg.Done()
|
f.wg.Done()
|
||||||
return errors.Join(errs...)
|
return errors.Join(errs...)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,73 +0,0 @@
|
|||||||
//go:build linux || darwin
|
|
||||||
|
|
||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
|
||||||
"github.com/slackhq/nebula/overlay/overlaytest"
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"github.com/slackhq/nebula/udp"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Test_emitStats_primesGauges covers issue #907: a Prometheus scrape that
|
|
||||||
// landed before the first ticker fire used to read 0 for the cert gauges.
|
|
||||||
// emitStats now primes the gauges before entering the ticker loop. We assert
|
|
||||||
// the gauge is zero before the first call and non-zero after.
|
|
||||||
func Test_emitStats_primesGauges(t *testing.T) {
|
|
||||||
defer metrics.DefaultRegistry.UnregisterAll()
|
|
||||||
|
|
||||||
l := test.NewLogger()
|
|
||||||
hostMap := newHostMap(l)
|
|
||||||
preferredRanges := []netip.Prefix{netip.MustParsePrefix("10.0.0.0/8")}
|
|
||||||
hostMap.preferredRanges.Store(&preferredRanges)
|
|
||||||
|
|
||||||
notAfter := time.Now().Add(time.Hour)
|
|
||||||
cs := &CertState{
|
|
||||||
initiatingVersion: cert.Version1,
|
|
||||||
privateKey: []byte{},
|
|
||||||
v1Cert: &dummyCert{version: cert.Version1, notAfter: notAfter},
|
|
||||||
v1Credential: nil,
|
|
||||||
}
|
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
|
||||||
ifce := &Interface{
|
|
||||||
hostMap: hostMap,
|
|
||||||
inside: &overlaytest.NoopTun{},
|
|
||||||
outside: &udp.NoopConn{},
|
|
||||||
firewall: &Firewall{Conntrack: &FirewallConntrack{Conns: map[firewall.Packet]*conn{}}},
|
|
||||||
lightHouse: lh,
|
|
||||||
pki: &PKI{},
|
|
||||||
handshakeManager: NewHandshakeManager(l, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
|
||||||
l: l,
|
|
||||||
// On linux, udp.NewUDPStatsEmitter indexes writers[0] and asserts to
|
|
||||||
// *udp.StdConn. A zero value works: getMemInfo sees a nil rawConn,
|
|
||||||
// returns an error, and the emitter falls through to a no-op.
|
|
||||||
writers: []udp.Conn{&udp.StdConn{}},
|
|
||||||
}
|
|
||||||
ifce.pki.cs.Store(cs)
|
|
||||||
|
|
||||||
ttlGauge := metrics.GetOrRegisterGauge("certificate.ttl_seconds", nil)
|
|
||||||
require.Zero(t, ttlGauge.Value(), "gauge should be zero before emitStats runs")
|
|
||||||
|
|
||||||
// Pre-cancel the context so emitStats returns after priming the gauges
|
|
||||||
// without ever reading from ticker.C. The one hour interval is just a
|
|
||||||
// belt-and-suspenders, the test does not expect the ticker to fire.
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
cancel()
|
|
||||||
ifce.emitStats(ctx, time.Hour)
|
|
||||||
|
|
||||||
ttl := ttlGauge.Value()
|
|
||||||
assert.Positive(t, ttl, "ttl gauge should be primed by emitStats before its first tick")
|
|
||||||
assert.LessOrEqual(t, ttl, int64(3600))
|
|
||||||
assert.Equal(t, int64(cert.Version1), metrics.GetOrRegisterGauge("certificate.initiating_version", nil).Value())
|
|
||||||
assert.Equal(t, int64(cert.Version1), metrics.GetOrRegisterGauge("certificate.max_version", nil).Value())
|
|
||||||
}
|
|
||||||
@@ -1,120 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/config"
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestReloadFirewall_CertUnsafeNetworksChanged verifies that reloadFirewall
|
|
||||||
// rebuilds the firewall when only the certificate's UnsafeNetworks have changed,
|
|
||||||
// even if the firewall section of the YAML has not.
|
|
||||||
func TestReloadFirewall_CertUnsafeNetworksChanged(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
|
|
||||||
vpnNet := netip.MustParsePrefix("10.0.0.1/24")
|
|
||||||
initialUnsafe := []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")}
|
|
||||||
|
|
||||||
// dummyCert avoids dragging the real signing pipeline into a unit test.
|
|
||||||
c1 := &dummyCert{
|
|
||||||
version: cert.Version2,
|
|
||||||
networks: []netip.Prefix{vpnNet},
|
|
||||||
unsafeNetworks: initialUnsafe,
|
|
||||||
}
|
|
||||||
pki := &PKI{}
|
|
||||||
pki.cs.Store(&CertState{v2Cert: c1, initiatingVersion: cert.Version2})
|
|
||||||
|
|
||||||
rawYAML := `firewall:
|
|
||||||
outbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
inbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
`
|
|
||||||
cfg := config.NewC(l)
|
|
||||||
require.NoError(t, cfg.LoadString(rawYAML))
|
|
||||||
|
|
||||||
fw, err := NewFirewallFromConfig(l, pki.getCertState(), cfg)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, initialUnsafe, fw.unsafeNetworks)
|
|
||||||
|
|
||||||
f := &Interface{
|
|
||||||
pki: pki,
|
|
||||||
firewall: fw,
|
|
||||||
l: l,
|
|
||||||
}
|
|
||||||
|
|
||||||
// Swap the cert with a different UnsafeNetworks set.
|
|
||||||
newUnsafe := []netip.Prefix{
|
|
||||||
netip.MustParsePrefix("198.51.100.0/24"),
|
|
||||||
netip.MustParsePrefix("203.0.113.0/24"),
|
|
||||||
}
|
|
||||||
c2 := &dummyCert{
|
|
||||||
version: cert.Version2,
|
|
||||||
networks: []netip.Prefix{vpnNet},
|
|
||||||
unsafeNetworks: newUnsafe,
|
|
||||||
}
|
|
||||||
pki.cs.Store(&CertState{v2Cert: c2, initiatingVersion: cert.Version2})
|
|
||||||
|
|
||||||
// Reload with the same YAML so HasChanged("firewall") reports false.
|
|
||||||
require.NoError(t, cfg.ReloadConfigString(rawYAML))
|
|
||||||
require.False(t, cfg.HasChanged("firewall"))
|
|
||||||
|
|
||||||
f.reloadFirewall(cfg)
|
|
||||||
|
|
||||||
assert.NotSame(t, fw, f.firewall, "firewall pointer should have been replaced")
|
|
||||||
assert.Equal(t, newUnsafe, f.firewall.unsafeNetworks)
|
|
||||||
assert.True(t, f.firewall.routableNetworks.Contains(netip.MustParseAddr("203.0.113.5")))
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestReloadFirewall_NoChange verifies that reloadFirewall is a no-op when
|
|
||||||
// neither the firewall config nor the cert's UnsafeNetworks have changed.
|
|
||||||
func TestReloadFirewall_NoChange(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
|
|
||||||
vpnNet := netip.MustParsePrefix("10.0.0.1/24")
|
|
||||||
unsafe := []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")}
|
|
||||||
|
|
||||||
c1 := &dummyCert{
|
|
||||||
version: cert.Version2,
|
|
||||||
networks: []netip.Prefix{vpnNet},
|
|
||||||
unsafeNetworks: unsafe,
|
|
||||||
}
|
|
||||||
pki := &PKI{}
|
|
||||||
pki.cs.Store(&CertState{v2Cert: c1, initiatingVersion: cert.Version2})
|
|
||||||
|
|
||||||
rawYAML := `firewall:
|
|
||||||
outbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
inbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
`
|
|
||||||
cfg := config.NewC(l)
|
|
||||||
require.NoError(t, cfg.LoadString(rawYAML))
|
|
||||||
|
|
||||||
fw, err := NewFirewallFromConfig(l, pki.getCertState(), cfg)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
f := &Interface{
|
|
||||||
pki: pki,
|
|
||||||
firewall: fw,
|
|
||||||
l: l,
|
|
||||||
}
|
|
||||||
|
|
||||||
require.NoError(t, cfg.ReloadConfigString(rawYAML))
|
|
||||||
f.reloadFirewall(cfg)
|
|
||||||
|
|
||||||
assert.Same(t, fw, f.firewall, "firewall should not have been replaced")
|
|
||||||
}
|
|
||||||
+8
-279
@@ -4,55 +4,27 @@ import (
|
|||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
|
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
"golang.org/x/net/ipv6"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// MaxIPv4RejectPacketSize is the largest IPv4 reject packet:
|
// Need 96 bytes for the largest reject packet:
|
||||||
// - 20 byte ipv4 header
|
// - 20 byte ipv4 header
|
||||||
// - 8 byte icmpv4 header
|
// - 8 byte icmpv4 header
|
||||||
// - 68 byte body (60 byte max orig ipv4 header + 8 byte orig icmpv4 header)
|
// - 68 byte body (60 byte max orig ipv4 header + 8 byte orig icmpv4 header)
|
||||||
maxIPv4RejectPacketSize = ipv4.HeaderLen + 8 + 60 + 8
|
MaxRejectPacketSize = ipv4.HeaderLen + 8 + 60 + 8
|
||||||
|
|
||||||
// MaxRejectPacketSize is sized for the largest possible reject packet (IPv6):
|
|
||||||
// - 40 byte ipv6 header
|
|
||||||
// - 8 byte icmpv6 header
|
|
||||||
// - up to 1000 byte body (original packet, possibly truncated. We want to stay
|
|
||||||
// under the MTU with Nebula overhead included)
|
|
||||||
maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000
|
|
||||||
|
|
||||||
MaxRejectPacketSize = maxIPv6RejectPacketSize
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func CreateRejectPacket(packet []byte, out []byte) []byte {
|
func CreateRejectPacket(packet []byte, out []byte) []byte {
|
||||||
if len(packet) < 1 {
|
if len(packet) < ipv4.HeaderLen || int(packet[0]>>4) != ipv4.Version {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
version := int(packet[0] >> 4)
|
|
||||||
switch version {
|
|
||||||
case ipv4.Version:
|
|
||||||
if len(packet) < ipv4.HeaderLen {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
// Do not send reject packets for non-first fragments
|
|
||||||
if packet[6]&0x1f != 0 || packet[7] != 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
switch packet[9] {
|
switch packet[9] {
|
||||||
case 6: // tcp
|
case 6: // tcp
|
||||||
return ipv4CreateRejectTCPPacket(packet, out)
|
return ipv4CreateRejectTCPPacket(packet, out)
|
||||||
default:
|
default:
|
||||||
return ipv4CreateRejectICMPPacket(packet, out)
|
return ipv4CreateRejectICMPPacket(packet, out)
|
||||||
}
|
}
|
||||||
case ipv6.Version:
|
|
||||||
if len(packet) < ipv6.HeaderLen {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return ipv6CreateRejectPacket(packet, out)
|
|
||||||
default:
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
|
func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
|
||||||
@@ -63,16 +35,11 @@ func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Do not generate ICMP errors in response to ICMP error packets
|
|
||||||
if packet[9] == 1 && len(packet) > ihl {
|
|
||||||
icmpType := packet[ihl]
|
|
||||||
if icmpType == 3 || icmpType == 4 || icmpType == 5 || icmpType == 11 || icmpType == 12 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ICMP reply includes original header and first 8 bytes of the packet
|
// ICMP reply includes original header and first 8 bytes of the packet
|
||||||
packetLen := min(len(packet), ihl+8)
|
packetLen := len(packet)
|
||||||
|
if packetLen > ihl+8 {
|
||||||
|
packetLen = ihl + 8
|
||||||
|
}
|
||||||
|
|
||||||
outLen := ipv4.HeaderLen + 8 + packetLen
|
outLen := ipv4.HeaderLen + 8 + packetLen
|
||||||
if outLen > cap(out) {
|
if outLen > cap(out) {
|
||||||
@@ -105,7 +72,7 @@ func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
|
|||||||
// ICMP Destination Unreachable
|
// ICMP Destination Unreachable
|
||||||
icmpOut := out[ipv4.HeaderLen:]
|
icmpOut := out[ipv4.HeaderLen:]
|
||||||
icmpOut[0] = 3 // type (Destination unreachable)
|
icmpOut[0] = 3 // type (Destination unreachable)
|
||||||
icmpOut[1] = 13 // code (Communication administratively prohibited)
|
icmpOut[1] = 3 // code (Port unreachable error)
|
||||||
icmpOut[2] = 0 // checksum
|
icmpOut[2] = 0 // checksum
|
||||||
icmpOut[3] = 0 // .
|
icmpOut[3] = 0 // .
|
||||||
icmpOut[4] = 0 // unused
|
icmpOut[4] = 0 // unused
|
||||||
@@ -198,193 +165,7 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte {
|
|||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
|
|
||||||
proto, offset, isFragment := ipv6FindUpperProtocol(packet)
|
|
||||||
if isFragment {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
switch proto {
|
|
||||||
case 6: // tcp
|
|
||||||
return ipv6CreateRejectTCPPacket(packet, out, offset)
|
|
||||||
default:
|
|
||||||
return ipv6CreateRejectICMPPacket(packet, out, proto, offset)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func ipv6CreateRejectICMPPacket(packet []byte, out []byte, proto uint8, offset int) []byte {
|
|
||||||
// Do not generate ICMPv6 errors in response to ICMPv6 error packets
|
|
||||||
if proto == 58 && len(packet) > offset {
|
|
||||||
icmpType := packet[offset]
|
|
||||||
if icmpType >= 1 && icmpType <= 4 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Include as much of the original packet as possible, up to 1000 bytes,
|
|
||||||
// so the response fits comfortably within any tunnel MTU.
|
|
||||||
packetLen := min(len(packet), 1000)
|
|
||||||
|
|
||||||
outLen := ipv6.HeaderLen + 8 + packetLen
|
|
||||||
if outLen > cap(out) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
out = out[:outLen]
|
|
||||||
|
|
||||||
// IPv6 header
|
|
||||||
ipHdr := out[0:ipv6.HeaderLen]
|
|
||||||
ipHdr[0] = ipv6.Version << 4 // version, traffic class (high bits)
|
|
||||||
ipHdr[1] = 0 // traffic class (low bits), flow label (high bits)
|
|
||||||
ipHdr[2] = 0 // flow label
|
|
||||||
ipHdr[3] = 0 // flow label
|
|
||||||
|
|
||||||
payloadLen := uint16(outLen - ipv6.HeaderLen)
|
|
||||||
binary.BigEndian.PutUint16(ipHdr[4:], payloadLen) // payload length
|
|
||||||
ipHdr[6] = 58 // next header (ICMPv6)
|
|
||||||
ipHdr[7] = 64 // hop limit
|
|
||||||
|
|
||||||
// Swap dest / src IPs (each 16 bytes, src at 8, dst at 24)
|
|
||||||
copy(ipHdr[8:24], packet[24:40])
|
|
||||||
copy(ipHdr[24:40], packet[8:24])
|
|
||||||
|
|
||||||
// ICMPv6 Destination Unreachable
|
|
||||||
icmpOut := out[ipv6.HeaderLen:]
|
|
||||||
icmpOut[0] = 1 // type (Destination Unreachable)
|
|
||||||
icmpOut[1] = 1 // code (Communication with destination administratively prohibited)
|
|
||||||
icmpOut[2] = 0 // checksum
|
|
||||||
icmpOut[3] = 0 // .
|
|
||||||
icmpOut[4] = 0 // unused
|
|
||||||
icmpOut[5] = 0 // .
|
|
||||||
icmpOut[6] = 0 // .
|
|
||||||
icmpOut[7] = 0 // .
|
|
||||||
|
|
||||||
copy(icmpOut[8:], packet[:packetLen])
|
|
||||||
|
|
||||||
// ICMPv6 checksum uses a pseudo-header
|
|
||||||
csum := ipv6PseudoheaderChecksum(ipHdr[8:24], ipHdr[24:40], 58, uint32(payloadLen))
|
|
||||||
binary.BigEndian.PutUint16(icmpOut[2:], tcpipChecksum(icmpOut, csum))
|
|
||||||
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte {
|
|
||||||
const tcpLen = 20
|
|
||||||
|
|
||||||
if len(packet) < offset+tcpLen {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
outLen := ipv6.HeaderLen + tcpLen
|
|
||||||
if outLen > cap(out) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
out = out[:outLen]
|
|
||||||
|
|
||||||
// IPv6 header
|
|
||||||
ipHdr := out[0:ipv6.HeaderLen]
|
|
||||||
ipHdr[0] = ipv6.Version << 4 // version, traffic class (high bits)
|
|
||||||
ipHdr[1] = 0 // traffic class (low bits), flow label (high bits)
|
|
||||||
ipHdr[2] = 0 // flow label
|
|
||||||
ipHdr[3] = 0 // flow label
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(ipHdr[4:], tcpLen) // payload length
|
|
||||||
ipHdr[6] = 6 // next header (TCP)
|
|
||||||
ipHdr[7] = 64 // hop limit
|
|
||||||
|
|
||||||
// Swap dest / src IPs
|
|
||||||
copy(ipHdr[8:24], packet[24:40])
|
|
||||||
copy(ipHdr[24:40], packet[8:24])
|
|
||||||
|
|
||||||
// TCP RST
|
|
||||||
tcpIn := packet[offset:]
|
|
||||||
var ackSeq, seq uint32
|
|
||||||
outFlags := byte(0b00000100) // RST
|
|
||||||
|
|
||||||
inAck := tcpIn[13]&0b00010000 != 0
|
|
||||||
if inAck {
|
|
||||||
seq = binary.BigEndian.Uint32(tcpIn[8:])
|
|
||||||
} else {
|
|
||||||
inSyn := uint32((tcpIn[13] & 0b00000010) >> 1)
|
|
||||||
inFin := uint32(tcpIn[13] & 0b00000001)
|
|
||||||
ackSeq = binary.BigEndian.Uint32(tcpIn[4:]) + inSyn + inFin + uint32(len(tcpIn)) - uint32(tcpIn[12]>>4)<<2
|
|
||||||
outFlags |= 0b00010000 // ACK
|
|
||||||
}
|
|
||||||
|
|
||||||
tcpOut := out[ipv6.HeaderLen:]
|
|
||||||
// Swap dest / src ports
|
|
||||||
copy(tcpOut[0:2], tcpIn[2:4])
|
|
||||||
copy(tcpOut[2:4], tcpIn[0:2])
|
|
||||||
binary.BigEndian.PutUint32(tcpOut[4:], seq)
|
|
||||||
binary.BigEndian.PutUint32(tcpOut[8:], ackSeq)
|
|
||||||
tcpOut[12] = (tcpLen >> 2) << 4 // data offset, reserved, NS
|
|
||||||
tcpOut[13] = outFlags // CWR, ECE, URG, ACK, PSH, RST, SYN, FIN
|
|
||||||
tcpOut[14] = 0 // window size
|
|
||||||
tcpOut[15] = 0 // .
|
|
||||||
tcpOut[16] = 0 // checksum
|
|
||||||
tcpOut[17] = 0 // .
|
|
||||||
tcpOut[18] = 0 // URG Pointer
|
|
||||||
tcpOut[19] = 0 // .
|
|
||||||
|
|
||||||
// Calculate checksum with IPv6 pseudo-header
|
|
||||||
csum := ipv6PseudoheaderChecksum(ipHdr[8:24], ipHdr[24:40], 6, tcpLen)
|
|
||||||
binary.BigEndian.PutUint16(tcpOut[16:], tcpipChecksum(tcpOut, csum))
|
|
||||||
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
|
|
||||||
nextHeader = packet[6]
|
|
||||||
offset = ipv6.HeaderLen
|
|
||||||
|
|
||||||
for {
|
|
||||||
switch nextHeader {
|
|
||||||
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
|
|
||||||
if len(packet) < offset+2 {
|
|
||||||
return nextHeader, offset, isFragment
|
|
||||||
}
|
|
||||||
nextHeader = packet[offset]
|
|
||||||
offset += (int(packet[offset+1]) + 1) << 3
|
|
||||||
|
|
||||||
case 44: // Fragment
|
|
||||||
if len(packet) < offset+8 {
|
|
||||||
return nextHeader, offset, isFragment
|
|
||||||
}
|
|
||||||
if packet[offset+2] != 0 || packet[offset+3]&0xf8 != 0 {
|
|
||||||
isFragment = true
|
|
||||||
}
|
|
||||||
nextHeader = packet[offset]
|
|
||||||
offset += 8
|
|
||||||
|
|
||||||
case 51: // AH
|
|
||||||
if len(packet) < offset+2 {
|
|
||||||
return nextHeader, offset, isFragment
|
|
||||||
}
|
|
||||||
nextHeader = packet[offset]
|
|
||||||
offset += (int(packet[offset+1]) + 2) << 2
|
|
||||||
|
|
||||||
default:
|
|
||||||
return nextHeader, offset, isFragment
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
||||||
if len(packet) < 1 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
switch packet[0] >> 4 {
|
|
||||||
case 4:
|
|
||||||
return createICMPv4EchoResponse(packet, out)
|
|
||||||
case 6:
|
|
||||||
return createICMPv6EchoResponse(packet, out)
|
|
||||||
default:
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func createICMPv4EchoResponse(packet, out []byte) []byte {
|
|
||||||
// Return early if this is not a simple ICMP Echo Request
|
// Return early if this is not a simple ICMP Echo Request
|
||||||
//TODO: make constants out of these
|
//TODO: make constants out of these
|
||||||
if !(len(packet) >= 28 && len(packet) <= 9001 && packet[0] == 0x45 && packet[9] == 0x01 && packet[20] == 0x08) {
|
if !(len(packet) >= 28 && len(packet) <= 9001 && packet[0] == 0x45 && packet[9] == 0x01 && packet[20] == 0x08) {
|
||||||
@@ -418,43 +199,6 @@ func createICMPv4EchoResponse(packet, out []byte) []byte {
|
|||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
func createICMPv6EchoResponse(packet, out []byte) []byte {
|
|
||||||
// IPv6 header (40 bytes) + ICMPv6 header (8 bytes minimum)
|
|
||||||
if len(packet) < ipv6.HeaderLen+8 || len(packet) > 9001 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Next Header must be ICMPv6 (58)
|
|
||||||
if packet[6] != 58 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// ICMPv6 type must be Echo Request (128)
|
|
||||||
if packet[ipv6.HeaderLen] != 128 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
out = out[:len(packet)]
|
|
||||||
copy(out, packet)
|
|
||||||
|
|
||||||
// Swap src/dst addresses (bytes 8-23 and 24-39)
|
|
||||||
copy(out[8:24], packet[24:40])
|
|
||||||
copy(out[24:40], packet[8:24])
|
|
||||||
|
|
||||||
// Change ICMPv6 type to Echo Reply (129)
|
|
||||||
icmp := out[ipv6.HeaderLen:]
|
|
||||||
icmp[0] = 129
|
|
||||||
icmp[2] = 0
|
|
||||||
icmp[3] = 0
|
|
||||||
|
|
||||||
// ICMPv6 checksum uses a pseudo-header with src, dst, length, and next header
|
|
||||||
payloadLen := uint32(len(icmp))
|
|
||||||
csum := ipv6PseudoheaderChecksum(out[8:24], out[24:40], 58, payloadLen)
|
|
||||||
binary.BigEndian.PutUint16(icmp[2:], tcpipChecksum(icmp, csum))
|
|
||||||
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
// calculates the TCP/IP checksum defined in rfc1071. The passed-in
|
// calculates the TCP/IP checksum defined in rfc1071. The passed-in
|
||||||
// csum is any initial checksum data that's already been computed.
|
// csum is any initial checksum data that's already been computed.
|
||||||
//
|
//
|
||||||
@@ -492,18 +236,3 @@ func ipv4PseudoheaderChecksum(src, dst []byte, proto, length uint32) (csum uint3
|
|||||||
csum += length >> 16
|
csum += length >> 16
|
||||||
return csum
|
return csum
|
||||||
}
|
}
|
||||||
|
|
||||||
// based on:
|
|
||||||
// - https://github.com/google/gopacket/blob/v1.1.19/layers/tcpip.go#L37-L48
|
|
||||||
func ipv6PseudoheaderChecksum(src, dst []byte, proto, length uint32) (csum uint32) {
|
|
||||||
for i := 0; i < 16; i += 2 {
|
|
||||||
csum += uint32(src[i]) << 8
|
|
||||||
csum += uint32(src[i+1])
|
|
||||||
csum += uint32(dst[i]) << 8
|
|
||||||
csum += uint32(dst[i+1])
|
|
||||||
}
|
|
||||||
csum += proto
|
|
||||||
csum += length & 0xffff
|
|
||||||
csum += length >> 16
|
|
||||||
return csum
|
|
||||||
}
|
|
||||||
|
|||||||
+1
-404
@@ -1,13 +1,11 @@
|
|||||||
package iputil
|
package iputil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
|
||||||
"net"
|
"net"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
"golang.org/x/net/ipv6"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func Test_CreateRejectPacket(t *testing.T) {
|
func Test_CreateRejectPacket(t *testing.T) {
|
||||||
@@ -45,7 +43,7 @@ func Test_CreateRejectPacket(t *testing.T) {
|
|||||||
}
|
}
|
||||||
b = append(b, []byte{0, 3, 0, 4, 0, 0, 0, 0}...)
|
b = append(b, []byte{0, 3, 0, 4, 0, 0, 0, 0}...)
|
||||||
|
|
||||||
expectedLen = maxIPv4RejectPacketSize
|
expectedLen = MaxRejectPacketSize
|
||||||
out = make([]byte, MaxRejectPacketSize)
|
out = make([]byte, MaxRejectPacketSize)
|
||||||
rejectPacket = CreateRejectPacket(b, out)
|
rejectPacket = CreateRejectPacket(b, out)
|
||||||
assert.NotNil(t, rejectPacket)
|
assert.NotNil(t, rejectPacket)
|
||||||
@@ -73,404 +71,3 @@ func Test_CreateRejectPacket(t *testing.T) {
|
|||||||
assert.NotNil(t, rejectPacket)
|
assert.NotNil(t, rejectPacket)
|
||||||
assert.Len(t, rejectPacket, expectedLen)
|
assert.Len(t, rejectPacket, expectedLen)
|
||||||
}
|
}
|
||||||
|
|
||||||
func Test_CreateRejectPacket_NoFragment(t *testing.T) {
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
|
|
||||||
// IPv4: non-zero fragment offset should not generate reject packet
|
|
||||||
h := ipv4.Header{
|
|
||||||
Len: 20,
|
|
||||||
Src: net.IPv4(10, 0, 0, 1),
|
|
||||||
Dst: net.IPv4(10, 0, 0, 2),
|
|
||||||
Protocol: 17, // UDP
|
|
||||||
}
|
|
||||||
b, err := h.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("h.Marshal: %v", err)
|
|
||||||
}
|
|
||||||
b = append(b, make([]byte, 8)...)
|
|
||||||
// Set fragment offset to non-zero (byte 6-7, offset in 8-byte units)
|
|
||||||
b[6] = 0x00
|
|
||||||
b[7] = 0x01
|
|
||||||
assert.Nil(t, CreateRejectPacket(b, out))
|
|
||||||
|
|
||||||
// MF flag with zero offset (first fragment) should still generate reject
|
|
||||||
b[6] = 0x20 // MF flag set
|
|
||||||
b[7] = 0x00
|
|
||||||
assert.NotNil(t, CreateRejectPacket(b, out))
|
|
||||||
|
|
||||||
// Non-fragment should still generate reject packet
|
|
||||||
b[6] = 0x00
|
|
||||||
b[7] = 0x00
|
|
||||||
assert.NotNil(t, CreateRejectPacket(b, out))
|
|
||||||
|
|
||||||
// DF flag only (not a fragment) should still generate reject packet
|
|
||||||
b[6] = 0x40
|
|
||||||
b[7] = 0x00
|
|
||||||
assert.NotNil(t, CreateRejectPacket(b, out))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_NoFragment(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
|
|
||||||
// IPv6 with Fragment header and non-zero offset should not generate reject
|
|
||||||
fragHeader := []byte{
|
|
||||||
17, // next header: UDP
|
|
||||||
0, // reserved
|
|
||||||
0, 9, // fragment offset=1 (shifted left 3), M=1
|
|
||||||
0, 0, 0, 1, // identification
|
|
||||||
}
|
|
||||||
udpPayload := make([]byte, 8)
|
|
||||||
payload := append(fragHeader, udpPayload...)
|
|
||||||
packet := makeIPv6Packet(src, dst, 44, payload) // next header 44 = Fragment
|
|
||||||
assert.Nil(t, CreateRejectPacket(packet, out))
|
|
||||||
|
|
||||||
// Fragment header with zero offset (first fragment) should still generate reject
|
|
||||||
fragHeader[2] = 0
|
|
||||||
fragHeader[3] = 1 // offset=0, M=1
|
|
||||||
payload = append(fragHeader, udpPayload...)
|
|
||||||
packet = makeIPv6Packet(src, dst, 44, payload)
|
|
||||||
assert.NotNil(t, CreateRejectPacket(packet, out))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacket_NoICMPError(t *testing.T) {
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
|
|
||||||
// ICMP error types should not generate reject packets
|
|
||||||
icmpErrorTypes := []byte{3, 4, 5, 11, 12}
|
|
||||||
for _, icmpType := range icmpErrorTypes {
|
|
||||||
h := ipv4.Header{
|
|
||||||
Len: 20,
|
|
||||||
Src: net.IPv4(10, 0, 0, 1),
|
|
||||||
Dst: net.IPv4(10, 0, 0, 2),
|
|
||||||
Protocol: 1, // ICMP
|
|
||||||
}
|
|
||||||
|
|
||||||
b, err := h.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("h.Marshal: %v", err)
|
|
||||||
}
|
|
||||||
b = append(b, icmpType, 0, 0, 0, 0, 0, 0, 0)
|
|
||||||
|
|
||||||
rejectPacket := CreateRejectPacket(b, out)
|
|
||||||
assert.Nil(t, rejectPacket, "ICMP type %d should not generate a reject packet", icmpType)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ICMP non-error types should still generate reject packets
|
|
||||||
icmpNonErrorTypes := []byte{0, 8, 13, 14}
|
|
||||||
for _, icmpType := range icmpNonErrorTypes {
|
|
||||||
h := ipv4.Header{
|
|
||||||
Len: 20,
|
|
||||||
Src: net.IPv4(10, 0, 0, 1),
|
|
||||||
Dst: net.IPv4(10, 0, 0, 2),
|
|
||||||
Protocol: 1, // ICMP
|
|
||||||
}
|
|
||||||
|
|
||||||
b, err := h.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("h.Marshal: %v", err)
|
|
||||||
}
|
|
||||||
b = append(b, icmpType, 0, 0, 0, 0, 0, 0, 0)
|
|
||||||
|
|
||||||
rejectPacket := CreateRejectPacket(b, out)
|
|
||||||
assert.NotNil(t, rejectPacket, "ICMP type %d should generate a reject packet", icmpType)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte {
|
|
||||||
b := make([]byte, ipv6.HeaderLen+len(payload))
|
|
||||||
b[0] = ipv6.Version << 4
|
|
||||||
binary.BigEndian.PutUint16(b[4:], uint16(len(payload)))
|
|
||||||
b[6] = nextHeader
|
|
||||||
b[7] = 64
|
|
||||||
copy(b[8:24], src.To16())
|
|
||||||
copy(b[24:40], dst.To16())
|
|
||||||
copy(b[ipv6.HeaderLen:], payload)
|
|
||||||
return b
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_ICMP(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
|
|
||||||
// Small UDP packet: entire original included in body
|
|
||||||
udpPayload := make([]byte, 20)
|
|
||||||
udpPayload[0] = 0x00 // src port high
|
|
||||||
udpPayload[1] = 0x50 // src port low (80)
|
|
||||||
udpPayload[2] = 0x01 // dst port high
|
|
||||||
udpPayload[3] = 0xBB // dst port low (443)
|
|
||||||
packet := makeIPv6Packet(src, dst, 17, udpPayload)
|
|
||||||
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.NotNil(t, rejectPacket)
|
|
||||||
|
|
||||||
// Small packet fits entirely: 40 (ipv6 hdr) + 8 (icmpv6 hdr) + 60 (original)
|
|
||||||
expectedLen := ipv6.HeaderLen + 8 + len(packet)
|
|
||||||
assert.Len(t, rejectPacket, expectedLen)
|
|
||||||
|
|
||||||
// Verify version
|
|
||||||
assert.Equal(t, byte(ipv6.Version<<4), rejectPacket[0]&0xf0)
|
|
||||||
// Verify next header is ICMPv6 (58)
|
|
||||||
assert.Equal(t, byte(58), rejectPacket[6])
|
|
||||||
// Verify src/dst are swapped
|
|
||||||
assert.Equal(t, dst.To16(), net.IP(rejectPacket[8:24]))
|
|
||||||
assert.Equal(t, src.To16(), net.IP(rejectPacket[24:40]))
|
|
||||||
// Verify ICMPv6 type=1 (Dest Unreachable), code=1 (Administratively prohibited)
|
|
||||||
assert.Equal(t, byte(1), rejectPacket[ipv6.HeaderLen])
|
|
||||||
assert.Equal(t, byte(1), rejectPacket[ipv6.HeaderLen+1])
|
|
||||||
// Verify entire original packet is included in body
|
|
||||||
assert.Equal(t, packet, rejectPacket[ipv6.HeaderLen+8:])
|
|
||||||
|
|
||||||
// Large packet: body is truncated to 1000 bytes
|
|
||||||
largePkt := makeIPv6Packet(src, dst, 17, make([]byte, 1200))
|
|
||||||
rejectPacket = CreateRejectPacket(largePkt, out)
|
|
||||||
assert.NotNil(t, rejectPacket)
|
|
||||||
assert.Len(t, rejectPacket, ipv6.HeaderLen+8+1000)
|
|
||||||
assert.Equal(t, largePkt[:1000], rejectPacket[ipv6.HeaderLen+8:])
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_TCP(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
|
|
||||||
// TCP SYN packet (next header 6)
|
|
||||||
tcpPayload := make([]byte, 20)
|
|
||||||
tcpPayload[0] = 0x00 // src port high
|
|
||||||
tcpPayload[1] = 0x50 // src port low (80)
|
|
||||||
tcpPayload[2] = 0x01 // dst port high
|
|
||||||
tcpPayload[3] = 0xBB // dst port low (443)
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[4:], 1000) // seq
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[8:], 0) // ack seq
|
|
||||||
tcpPayload[12] = (20 >> 2) << 4 // data offset
|
|
||||||
tcpPayload[13] = 0b00000010 // SYN flag
|
|
||||||
|
|
||||||
packet := makeIPv6Packet(src, dst, 6, tcpPayload)
|
|
||||||
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.NotNil(t, rejectPacket)
|
|
||||||
|
|
||||||
// Expected: 40 (ipv6 hdr) + 20 (tcp RST)
|
|
||||||
expectedLen := ipv6.HeaderLen + 20
|
|
||||||
assert.Len(t, rejectPacket, expectedLen)
|
|
||||||
|
|
||||||
// Verify version
|
|
||||||
assert.Equal(t, byte(ipv6.Version<<4), rejectPacket[0]&0xf0)
|
|
||||||
// Verify next header is TCP (6)
|
|
||||||
assert.Equal(t, byte(6), rejectPacket[6])
|
|
||||||
// Verify src/dst are swapped
|
|
||||||
assert.Equal(t, dst.To16(), net.IP(rejectPacket[8:24]))
|
|
||||||
assert.Equal(t, src.To16(), net.IP(rejectPacket[24:40]))
|
|
||||||
// Verify ports are swapped
|
|
||||||
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
|
||||||
assert.Equal(t, uint16(443), binary.BigEndian.Uint16(tcpOut[0:2]))
|
|
||||||
assert.Equal(t, uint16(80), binary.BigEndian.Uint16(tcpOut[2:4]))
|
|
||||||
// RST+ACK flags (since input was SYN without ACK)
|
|
||||||
assert.Equal(t, byte(0b00010100), tcpOut[13])
|
|
||||||
// ack_seq = original seq (1000) + SYN (1) + FIN (0) + segment data (0)
|
|
||||||
assert.Equal(t, uint32(1001), binary.BigEndian.Uint32(tcpOut[8:]))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_TCPWithACK(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
|
|
||||||
// TCP packet with ACK set
|
|
||||||
tcpPayload := make([]byte, 20)
|
|
||||||
tcpPayload[0] = 0x00
|
|
||||||
tcpPayload[1] = 0x50
|
|
||||||
tcpPayload[2] = 0x01
|
|
||||||
tcpPayload[3] = 0xBB
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[4:], 1000) // seq
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[8:], 2000) // ack seq
|
|
||||||
tcpPayload[12] = (20 >> 2) << 4 // data offset
|
|
||||||
tcpPayload[13] = 0b00010000 // ACK flag
|
|
||||||
|
|
||||||
packet := makeIPv6Packet(src, dst, 6, tcpPayload)
|
|
||||||
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.NotNil(t, rejectPacket)
|
|
||||||
|
|
||||||
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
|
||||||
// RST only (no ACK) since input had ACK
|
|
||||||
assert.Equal(t, byte(0b00000100), tcpOut[13])
|
|
||||||
// seq = original ack_seq
|
|
||||||
assert.Equal(t, uint32(2000), binary.BigEndian.Uint32(tcpOut[4:]))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_NoICMPError(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
|
|
||||||
// ICMPv6 error types (1-4) should not generate reject packets
|
|
||||||
for icmpType := byte(1); icmpType <= 4; icmpType++ {
|
|
||||||
payload := make([]byte, 8)
|
|
||||||
payload[0] = icmpType
|
|
||||||
packet := makeIPv6Packet(src, dst, 58, payload)
|
|
||||||
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.Nil(t, rejectPacket, "ICMPv6 type %d should not generate a reject packet", icmpType)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ICMPv6 non-error types should still generate reject packets
|
|
||||||
nonErrorTypes := []byte{128, 129, 133, 134}
|
|
||||||
for _, icmpType := range nonErrorTypes {
|
|
||||||
payload := make([]byte, 8)
|
|
||||||
payload[0] = icmpType
|
|
||||||
packet := makeIPv6Packet(src, dst, 58, payload)
|
|
||||||
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.NotNil(t, rejectPacket, "ICMPv6 type %d should generate a reject packet", icmpType)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_TooShort(t *testing.T) {
|
|
||||||
// Packet too short to be valid IPv6
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
assert.Nil(t, CreateRejectPacket([]byte{0x60}, out))
|
|
||||||
assert.Nil(t, CreateRejectPacket(make([]byte, 39), out))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_ExtensionHeaders(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
|
|
||||||
// IPv6 + Hop-by-Hop extension header + TCP
|
|
||||||
hopByHop := []byte{
|
|
||||||
6, // next header: TCP
|
|
||||||
0, // length (8 bytes total)
|
|
||||||
0, 0, // padding
|
|
||||||
0, 0, 0, 0,
|
|
||||||
}
|
|
||||||
tcpPayload := make([]byte, 20)
|
|
||||||
tcpPayload[0] = 0x00
|
|
||||||
tcpPayload[1] = 0x50
|
|
||||||
tcpPayload[2] = 0x01
|
|
||||||
tcpPayload[3] = 0xBB
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[4:], 1000)
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[8:], 2000)
|
|
||||||
tcpPayload[12] = (20 >> 2) << 4
|
|
||||||
tcpPayload[13] = 0b00010000 // ACK
|
|
||||||
|
|
||||||
payload := append(hopByHop, tcpPayload...)
|
|
||||||
packet := makeIPv6Packet(src, dst, 0, payload) // next header 0 = Hop-by-Hop
|
|
||||||
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.NotNil(t, rejectPacket)
|
|
||||||
|
|
||||||
// Should produce TCP RST
|
|
||||||
expectedLen := ipv6.HeaderLen + 20
|
|
||||||
assert.Len(t, rejectPacket, expectedLen)
|
|
||||||
assert.Equal(t, byte(6), rejectPacket[6]) // next header is TCP
|
|
||||||
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
|
||||||
assert.Equal(t, byte(0b00000100), tcpOut[13]) // RST only
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateICMPEchoResponse_IPv4(t *testing.T) {
|
|
||||||
// Build a simple IPv4 ICMP Echo Request
|
|
||||||
packet := make([]byte, 28)
|
|
||||||
packet[0] = 0x45 // version 4, IHL 5
|
|
||||||
binary.BigEndian.PutUint16(packet[2:], uint16(28)) // total length
|
|
||||||
packet[8] = 64 // TTL
|
|
||||||
packet[9] = 1 // protocol ICMP
|
|
||||||
copy(packet[12:16], net.IPv4(10, 0, 0, 1).To4()) // src
|
|
||||||
copy(packet[16:20], net.IPv4(10, 0, 0, 2).To4()) // dst
|
|
||||||
packet[20] = 8 // ICMP Echo Request
|
|
||||||
|
|
||||||
out := make([]byte, len(packet))
|
|
||||||
result := CreateICMPEchoResponse(packet, out)
|
|
||||||
assert.NotNil(t, result)
|
|
||||||
assert.Equal(t, byte(0x45), result[0])
|
|
||||||
// src/dst swapped
|
|
||||||
assert.Equal(t, net.IPv4(10, 0, 0, 2).To4(), net.IP(result[12:16]))
|
|
||||||
assert.Equal(t, net.IPv4(10, 0, 0, 1).To4(), net.IP(result[16:20]))
|
|
||||||
// ICMP Echo Reply
|
|
||||||
assert.Equal(t, byte(0), result[20])
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateICMPEchoResponse_IPv6(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1").To16()
|
|
||||||
dst := net.ParseIP("fd00::2").To16()
|
|
||||||
|
|
||||||
// Build an IPv6 ICMPv6 Echo Request packet
|
|
||||||
// IPv6 header (40 bytes) + ICMPv6 (8 bytes)
|
|
||||||
packet := make([]byte, 48)
|
|
||||||
packet[0] = 0x60 // version 6
|
|
||||||
payloadLen := uint16(8) // ICMPv6 header only
|
|
||||||
binary.BigEndian.PutUint16(packet[4:], payloadLen)
|
|
||||||
packet[6] = 58 // Next Header: ICMPv6
|
|
||||||
packet[7] = 64 // Hop Limit
|
|
||||||
copy(packet[8:24], src) // src address
|
|
||||||
copy(packet[24:40], dst) // dst address
|
|
||||||
|
|
||||||
// ICMPv6 Echo Request
|
|
||||||
icmp := packet[40:]
|
|
||||||
icmp[0] = 128 // type: Echo Request
|
|
||||||
icmp[1] = 0 // code
|
|
||||||
binary.BigEndian.PutUint16(icmp[4:], 1) // identifier
|
|
||||||
binary.BigEndian.PutUint16(icmp[6:], 1) // sequence number
|
|
||||||
|
|
||||||
// Compute correct checksum for the request
|
|
||||||
csum := ipv6PseudoheaderChecksum(src, dst, 58, uint32(payloadLen))
|
|
||||||
binary.BigEndian.PutUint16(icmp[2:], tcpipChecksum(icmp, csum))
|
|
||||||
|
|
||||||
out := make([]byte, len(packet))
|
|
||||||
result := CreateICMPEchoResponse(packet, out)
|
|
||||||
assert.NotNil(t, result)
|
|
||||||
|
|
||||||
// Version should still be 6
|
|
||||||
assert.Equal(t, byte(6), result[0]>>4)
|
|
||||||
// src/dst swapped
|
|
||||||
assert.Equal(t, dst, net.IP(result[8:24]))
|
|
||||||
assert.Equal(t, src, net.IP(result[24:40]))
|
|
||||||
// ICMPv6 Echo Reply type
|
|
||||||
assert.Equal(t, byte(129), result[40])
|
|
||||||
|
|
||||||
// Verify checksum is valid (tcpipChecksum returns 0 when data+checksum is correct)
|
|
||||||
respIcmp := result[40:]
|
|
||||||
verifyCsum := ipv6PseudoheaderChecksum(result[8:24], result[24:40], 58, uint32(payloadLen))
|
|
||||||
assert.Equal(t, uint16(0), tcpipChecksum(respIcmp, verifyCsum))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateICMPEchoResponse_IPv6_NotEchoRequest(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1").To16()
|
|
||||||
dst := net.ParseIP("fd00::2").To16()
|
|
||||||
|
|
||||||
packet := make([]byte, 48)
|
|
||||||
packet[0] = 0x60
|
|
||||||
binary.BigEndian.PutUint16(packet[4:], 8)
|
|
||||||
packet[6] = 58
|
|
||||||
packet[7] = 64
|
|
||||||
copy(packet[8:24], src)
|
|
||||||
copy(packet[24:40], dst)
|
|
||||||
|
|
||||||
// ICMPv6 type 1 (Destination Unreachable) - not Echo Request
|
|
||||||
packet[40] = 1
|
|
||||||
|
|
||||||
out := make([]byte, len(packet))
|
|
||||||
result := CreateICMPEchoResponse(packet, out)
|
|
||||||
assert.Nil(t, result)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1").To16()
|
|
||||||
dst := net.ParseIP("fd00::2").To16()
|
|
||||||
|
|
||||||
packet := make([]byte, 48)
|
|
||||||
packet[0] = 0x60
|
|
||||||
binary.BigEndian.PutUint16(packet[4:], 8)
|
|
||||||
packet[6] = 6 // TCP, not ICMPv6
|
|
||||||
packet[7] = 64
|
|
||||||
copy(packet[8:24], src)
|
|
||||||
copy(packet[24:40], dst)
|
|
||||||
|
|
||||||
out := make([]byte, len(packet))
|
|
||||||
result := CreateICMPEchoResponse(packet, out)
|
|
||||||
assert.Nil(t, result)
|
|
||||||
}
|
|
||||||
|
|||||||
+50
-43
@@ -15,6 +15,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
@@ -34,12 +35,9 @@ type LightHouse struct {
|
|||||||
|
|
||||||
myVpnNetworks []netip.Prefix
|
myVpnNetworks []netip.Prefix
|
||||||
myVpnNetworksTable *bart.Lite
|
myVpnNetworksTable *bart.Lite
|
||||||
|
punchConn udp.Conn
|
||||||
punchy *Punchy
|
punchy *Punchy
|
||||||
|
|
||||||
// localAddrsFn enumerates the underlay addresses we advertise. It is a field so tests can supply simulated
|
|
||||||
// addresses rather than whatever this machine's NICs happen to be. Set it before Start.
|
|
||||||
localAddrsFn func(*LocalAllowList) []netip.Addr
|
|
||||||
|
|
||||||
// Local cache of answers from light houses
|
// Local cache of answers from light houses
|
||||||
// map of vpn addr to answers
|
// map of vpn addr to answers
|
||||||
addrMap map[netip.Addr]*RemoteList
|
addrMap map[netip.Addr]*RemoteList
|
||||||
@@ -78,6 +76,7 @@ type LightHouse struct {
|
|||||||
calculatedRemotes atomic.Pointer[bart.Table[[]*calculatedRemote]] // Maps VpnAddr to []*calculatedRemote
|
calculatedRemotes atomic.Pointer[bart.Table[[]*calculatedRemote]] // Maps VpnAddr to []*calculatedRemote
|
||||||
|
|
||||||
metrics *MessageMetrics
|
metrics *MessageMetrics
|
||||||
|
metricHolepunchTx metrics.Counter
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -106,15 +105,12 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
|||||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||||
addrMap: make(map[netip.Addr]*RemoteList),
|
addrMap: make(map[netip.Addr]*RemoteList),
|
||||||
nebulaPort: nebulaPort,
|
nebulaPort: nebulaPort,
|
||||||
|
punchConn: pc,
|
||||||
punchy: p,
|
punchy: p,
|
||||||
updateTrigger: make(chan struct{}, 1),
|
updateTrigger: make(chan struct{}, 1),
|
||||||
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
h.localAddrsFn = func(al *LocalAllowList) []netip.Addr {
|
|
||||||
return localAddrs(h.l, al)
|
|
||||||
}
|
|
||||||
|
|
||||||
lighthouses := make([]netip.Addr, 0)
|
lighthouses := make([]netip.Addr, 0)
|
||||||
h.lighthouses.Store(&lighthouses)
|
h.lighthouses.Store(&lighthouses)
|
||||||
staticList := make(map[netip.Addr]struct{})
|
staticList := make(map[netip.Addr]struct{})
|
||||||
@@ -122,6 +118,9 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
|||||||
|
|
||||||
if c.GetBool("stats.lighthouse_metrics", false) {
|
if c.GetBool("stats.lighthouse_metrics", false) {
|
||||||
h.metrics = newLighthouseMetrics()
|
h.metrics = newLighthouseMetrics()
|
||||||
|
h.metricHolepunchTx = metrics.GetOrRegisterCounter("messages.tx.holepunch", nil)
|
||||||
|
} else {
|
||||||
|
h.metricHolepunchTx = metrics.NilCounter{}
|
||||||
}
|
}
|
||||||
|
|
||||||
err := h.reload(c, true)
|
err := h.reload(c, true)
|
||||||
@@ -280,18 +279,16 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
|
|||||||
//NOTE: many things will get much simpler when we combine static_host_map and lighthouse.hosts in config
|
//NOTE: many things will get much simpler when we combine static_host_map and lighthouse.hosts in config
|
||||||
if initial || c.HasChanged("static_host_map") || c.HasChanged("static_map.cadence") || c.HasChanged("static_map.network") || c.HasChanged("static_map.lookup_timeout") {
|
if initial || c.HasChanged("static_host_map") || c.HasChanged("static_map.cadence") || c.HasChanged("static_map.network") || c.HasChanged("static_map.lookup_timeout") {
|
||||||
// Clean up. Entries still in the static_host_map will be re-built.
|
// Clean up. Entries still in the static_host_map will be re-built.
|
||||||
ourselves := lh.myVpnNetworks[0].Addr()
|
// Entries no longer present must have their (possible) background DNS goroutines stopped.
|
||||||
oldStaticList := lh.staticList.Load()
|
if existingStaticList := lh.staticList.Load(); existingStaticList != nil {
|
||||||
if oldStaticList != nil {
|
|
||||||
lh.RLock()
|
lh.RLock()
|
||||||
for staticVpnAddr := range *oldStaticList {
|
for staticVpnAddr := range *existingStaticList {
|
||||||
if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil {
|
if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil {
|
||||||
am.ResetForOwner(ourselves)
|
am.hr.Cancel()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
lh.RUnlock()
|
lh.RUnlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build a new list based on current config.
|
// Build a new list based on current config.
|
||||||
staticList := make(map[netip.Addr]struct{})
|
staticList := make(map[netip.Addr]struct{})
|
||||||
err := lh.loadStaticMap(c, staticList)
|
err := lh.loadStaticMap(c, staticList)
|
||||||
@@ -299,21 +296,6 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// For entries removed from static_host_map, stop the DNS goroutine and drop the cached addrs.
|
|
||||||
// All addrs must come from the lighthouses now that it's no longer a static host.
|
|
||||||
if oldStaticList != nil {
|
|
||||||
lh.RLock()
|
|
||||||
for staticVpnAddr := range *oldStaticList {
|
|
||||||
if _, stillStatic := staticList[staticVpnAddr]; stillStatic {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil {
|
|
||||||
am.ClearHostnameResults()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
lh.RUnlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
lh.staticList.Store(&staticList)
|
lh.staticList.Store(&staticList)
|
||||||
if !initial {
|
if !initial {
|
||||||
if c.HasChanged("static_host_map") {
|
if c.HasChanged("static_host_map") {
|
||||||
@@ -926,7 +908,7 @@ func (lh *LightHouse) SendUpdate() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
lal := lh.GetLocalAllowList()
|
lal := lh.GetLocalAllowList()
|
||||||
for _, e := range lh.localAddrsFn(lal) {
|
for _, e := range localAddrs(lh.l, lal) {
|
||||||
if lh.myVpnNetworksTable.Contains(e) {
|
if lh.myVpnNetworksTable.Contains(e) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -1424,31 +1406,58 @@ func (lhh *LightHouseHandler) handleHostPunchNotification(n *NebulaMeta, fromVpn
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
empty := []byte{0}
|
||||||
|
punch := func(vpnPeer netip.AddrPort, logVpnAddr netip.Addr) {
|
||||||
|
if !vpnPeer.IsValid() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
time.Sleep(lhh.lh.punchy.GetDelay())
|
||||||
|
lhh.lh.metricHolepunchTx.Inc(1)
|
||||||
|
lhh.lh.punchConn.WriteTo(empty, vpnPeer)
|
||||||
|
}()
|
||||||
|
|
||||||
|
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
lhh.l.Debug("Punching",
|
||||||
|
"vpnPeer", vpnPeer,
|
||||||
|
"logVpnAddr", logVpnAddr,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
remoteAllowList := lhh.lh.GetRemoteAllowList()
|
remoteAllowList := lhh.lh.GetRemoteAllowList()
|
||||||
for _, a := range n.Details.V4AddrPorts {
|
for _, a := range n.Details.V4AddrPorts {
|
||||||
if a == nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
b := protoV4AddrPortToNetAddrPort(a)
|
b := protoV4AddrPortToNetAddrPort(a)
|
||||||
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
|
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
|
||||||
lhh.lh.punchy.Schedule(b, detailsVpnAddr)
|
punch(b, detailsVpnAddr)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, a := range n.Details.V6AddrPorts {
|
for _, a := range n.Details.V6AddrPorts {
|
||||||
if a == nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
b := protoV6AddrPortToNetAddrPort(a)
|
b := protoV6AddrPortToNetAddrPort(a)
|
||||||
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
|
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
|
||||||
lhh.lh.punchy.Schedule(b, detailsVpnAddr)
|
punch(b, detailsVpnAddr)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// This sends a nebula test packet to the host trying to contact us. In the case
|
// This sends a nebula test packet to the host trying to contact us. In the case
|
||||||
// of a double nat or other difficult scenario, this may help establish
|
// of a double nat or other difficult scenario, this may help establish
|
||||||
// a tunnel. ScheduleRespond is a no-op when punchy.respond is disabled.
|
// a tunnel.
|
||||||
lhh.lh.punchy.ScheduleRespond(detailsVpnAddr)
|
if lhh.lh.punchy.GetRespond() {
|
||||||
|
go func() {
|
||||||
|
time.Sleep(lhh.lh.punchy.GetRespondDelay())
|
||||||
|
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
lhh.l.Debug("Sending a nebula test packet",
|
||||||
|
"vpnAddr", detailsVpnAddr,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
//NOTE: we have to allocate a new output buffer here since we are spawning a new goroutine
|
||||||
|
// for each punchBack packet. We should move this into a timerwheel or a single goroutine
|
||||||
|
// managed by a channel.
|
||||||
|
w.SendMessageToVpnAddr(header.Test, header.TestRequest, detailsVpnAddr, []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
||||||
|
}()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func protoAddrToNetAddr(addr *Addr) netip.Addr {
|
func protoAddrToNetAddr(addr *Addr) netip.Addr {
|
||||||
@@ -1468,7 +1477,7 @@ func protoV6AddrPortToNetAddrPort(ap *V6AddrPort) netip.AddrPort {
|
|||||||
b := [16]byte{}
|
b := [16]byte{}
|
||||||
binary.BigEndian.PutUint64(b[:8], ap.Hi)
|
binary.BigEndian.PutUint64(b[:8], ap.Hi)
|
||||||
binary.BigEndian.PutUint64(b[8:], ap.Lo)
|
binary.BigEndian.PutUint64(b[8:], ap.Lo)
|
||||||
return netip.AddrPortFrom(netip.AddrFrom16(b).Unmap(), uint16(ap.Port))
|
return netip.AddrPortFrom(netip.AddrFrom16(b), uint16(ap.Port))
|
||||||
}
|
}
|
||||||
|
|
||||||
func netAddrToProtoAddr(addr netip.Addr) *Addr {
|
func netAddrToProtoAddr(addr netip.Addr) *Addr {
|
||||||
@@ -1508,11 +1517,9 @@ func (d *NebulaMetaDetails) GetRelays() []netip.Addr {
|
|||||||
|
|
||||||
if len(d.RelayVpnAddrs) > 0 {
|
if len(d.RelayVpnAddrs) > 0 {
|
||||||
for _, r := range d.RelayVpnAddrs {
|
for _, r := range d.RelayVpnAddrs {
|
||||||
if r != nil {
|
|
||||||
relays = append(relays, protoAddrToNetAddr(r))
|
relays = append(relays, protoAddrToNetAddr(r))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
|
||||||
return relays
|
return relays
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -303,132 +303,6 @@ func TestLighthouse_reload(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestLighthouse_reloadStaticHostMap verifies that reloading static_host_map applies the new
|
|
||||||
// config rather than appending to it. See issue #718.
|
|
||||||
func TestLighthouse_reloadStaticHostMap(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
c := config.NewC(l)
|
|
||||||
c.Settings["lighthouse"] = map[string]any{"am_lighthouse": true}
|
|
||||||
c.Settings["listen"] = map[string]any{"port": 4242}
|
|
||||||
c.Settings["static_host_map"] = map[string]any{
|
|
||||||
"10.128.0.2": []any{"1.1.1.1:4242"},
|
|
||||||
}
|
|
||||||
|
|
||||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
|
||||||
nt := new(bart.Lite)
|
|
||||||
nt.Insert(myVpnNet)
|
|
||||||
cs := &CertState{
|
|
||||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
|
||||||
myVpnNetworksTable: nt,
|
|
||||||
}
|
|
||||||
|
|
||||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
staticHost := netip.MustParseAddr("10.128.0.2")
|
|
||||||
otherHost := netip.MustParseAddr("10.128.0.3")
|
|
||||||
|
|
||||||
// Capture the RemoteList pointer up front; an in-flight handshake would hold the same one
|
|
||||||
// on hostinfo.remotes, so it must reflect every reload below.
|
|
||||||
pinned := lh.Query(staticHost)
|
|
||||||
require.NotNil(t, pinned)
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, pinned.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
// Replace the remote address. The new address should be the only entry.
|
|
||||||
nc := map[string]any{
|
|
||||||
"static_host_map": map[string]any{
|
|
||||||
"10.128.0.2": []any{"2.2.2.2:4242"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
rc, err := yaml.Marshal(nc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
rl := lh.Query(staticHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Same(t, pinned, rl, "RemoteList pointer must stay stable so in-flight handshakes pick up the change")
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("2.2.2.2:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
// Reload back to the original IP. Mirrors the round-trip in issue #718 step 6-8 where
|
|
||||||
// the buggy reload produced [1.1.1.1, 2.2.2.2, 1.1.1.1] instead of [1.1.1.1].
|
|
||||||
nc = map[string]any{
|
|
||||||
"static_host_map": map[string]any{
|
|
||||||
"10.128.0.2": []any{"1.1.1.1:4242"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
rc, err = yaml.Marshal(nc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
rl = lh.Query(staticHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Same(t, pinned, rl)
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
// Reload with the same config. An unchanged entry must not duplicate.
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
rl = lh.Query(staticHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Same(t, pinned, rl)
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
// Switch back to 2.2.2.2 so the rest of the test continues against a known address.
|
|
||||||
nc = map[string]any{
|
|
||||||
"static_host_map": map[string]any{
|
|
||||||
"10.128.0.2": []any{"2.2.2.2:4242"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
rc, err = yaml.Marshal(nc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
// Add a second host alongside the first. Both should be present, neither duplicated.
|
|
||||||
nc = map[string]any{
|
|
||||||
"static_host_map": map[string]any{
|
|
||||||
"10.128.0.2": []any{"2.2.2.2:4242"},
|
|
||||||
"10.128.0.3": []any{"3.3.3.3:4242"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
rc, err = yaml.Marshal(nc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
rl = lh.Query(staticHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Same(t, pinned, rl, "adding a sibling entry must not displace the existing RemoteList")
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("2.2.2.2:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
rl = lh.Query(otherHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
// Drop the first host entirely. The vpnAddr is no longer marked static, our owner
|
|
||||||
// contribution is cleared, but the addrMap entry stays in place so non-static cache
|
|
||||||
// data (from lighthouse queries) on the same RemoteList isn't lost. In-flight handshakes
|
|
||||||
// that already had the pointer see an empty address list rather than retrying stale ones.
|
|
||||||
nc = map[string]any{
|
|
||||||
"static_host_map": map[string]any{
|
|
||||||
"10.128.0.3": []any{"3.3.3.3:4242"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
rc, err = yaml.Marshal(nc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
_, isStatic := lh.GetStaticHostList()[staticHost]
|
|
||||||
assert.False(t, isStatic)
|
|
||||||
|
|
||||||
rl = lh.Query(staticHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Same(t, pinned, rl)
|
|
||||||
assert.Empty(t, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
rl = lh.Query(otherHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
}
|
|
||||||
|
|
||||||
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
||||||
req := &NebulaMeta{
|
req := &NebulaMeta{
|
||||||
Type: NebulaMeta_HostQuery,
|
Type: NebulaMeta_HostQuery,
|
||||||
|
|||||||
@@ -55,7 +55,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
}
|
}
|
||||||
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
|
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
|
||||||
|
|
||||||
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"))
|
ssh, err := sshd.NewSSHServer(l.With("subsystem", "sshd"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
|
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
|
||||||
}
|
}
|
||||||
@@ -130,17 +130,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
udpConns := make([]udp.Conn, routines)
|
udpConns := make([]udp.Conn, routines)
|
||||||
port := c.GetInt("listen.port", 0)
|
port := c.GetInt("listen.port", 0)
|
||||||
|
|
||||||
// Callers get no handle to these until the Control is returned, release them on any error.
|
|
||||||
defer func() {
|
|
||||||
if reterr != nil {
|
|
||||||
for _, u := range udpConns {
|
|
||||||
if u != nil {
|
|
||||||
_ = u.Close()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
if !configTest {
|
if !configTest {
|
||||||
rawListenHost := c.GetString("listen.host", "0.0.0.0")
|
rawListenHost := c.GetString("listen.host", "0.0.0.0")
|
||||||
var listenHost netip.Addr
|
var listenHost netip.Addr
|
||||||
@@ -181,7 +170,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
}
|
}
|
||||||
|
|
||||||
hostMap := NewHostMapFromConfig(l, c)
|
hostMap := NewHostMapFromConfig(l, c)
|
||||||
punchy := NewPunchyFromConfig(l, c, udpConns[0])
|
punchy := NewPunchyFromConfig(l, c)
|
||||||
connManager := newConnectionManagerFromConfig(l, c, hostMap, punchy)
|
connManager := newConnectionManagerFromConfig(l, c, hostMap, punchy)
|
||||||
lightHouse, err := NewLightHouseFromConfig(ctx, l, c, pki.getCertState(), udpConns[0], punchy)
|
lightHouse, err := NewLightHouseFromConfig(ctx, l, c, pki.getCertState(), udpConns[0], punchy)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -195,17 +184,21 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
messageMetrics = newMessageMetricsOnlyRecvError()
|
messageMetrics = newMessageMetricsOnlyRecvError()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
useRelays := c.GetBool("relay.use_relays", DefaultUseRelays) && !c.GetBool("relay.am_relay", false)
|
||||||
|
|
||||||
handshakeConfig := HandshakeConfig{
|
handshakeConfig := HandshakeConfig{
|
||||||
tryInterval: c.GetDuration("handshakes.try_interval", DefaultHandshakeTryInterval),
|
tryInterval: c.GetDuration("handshakes.try_interval", DefaultHandshakeTryInterval),
|
||||||
retries: int64(c.GetInt("handshakes.retries", DefaultHandshakeRetries)),
|
retries: int64(c.GetInt("handshakes.retries", DefaultHandshakeRetries)),
|
||||||
triggerBuffer: c.GetInt("handshakes.trigger_buffer", DefaultHandshakeTriggerBuffer),
|
triggerBuffer: c.GetInt("handshakes.trigger_buffer", DefaultHandshakeTriggerBuffer),
|
||||||
|
useRelays: useRelays,
|
||||||
|
|
||||||
messageMetrics: messageMetrics,
|
messageMetrics: messageMetrics,
|
||||||
}
|
}
|
||||||
|
|
||||||
handshakeManager := NewHandshakeManager(l, hostMap, lightHouse, udpConns[0], handshakeConfig)
|
handshakeManager := NewHandshakeManager(l, hostMap, lightHouse, udpConns[0], handshakeConfig)
|
||||||
lightHouse.handshakeTrigger = handshakeManager.trigger
|
lightHouse.handshakeTrigger = handshakeManager.trigger
|
||||||
|
|
||||||
ds, err := newDnsServerFromConfig(ctx, l, pki, hostMap, c)
|
ds, err := newDnsServerFromConfig(ctx, l, pki.getCertState(), hostMap, c)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
l.Warn("Failed to start DNS responder", "error", err)
|
l.Warn("Failed to start DNS responder", "error", err)
|
||||||
}
|
}
|
||||||
@@ -251,8 +244,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
|
|
||||||
handshakeManager.f = ifce
|
handshakeManager.f = ifce
|
||||||
go handshakeManager.Run(ctx)
|
go handshakeManager.Run(ctx)
|
||||||
|
|
||||||
punchy.Start(ctx, ifce, hostMap, lightHouse)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
stats, err := newStatsServerFromConfig(ctx, l, c, buildVersion, configTest)
|
stats, err := newStatsServerFromConfig(ctx, l, c, buildVersion, configTest)
|
||||||
@@ -268,8 +259,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
|
|
||||||
attachCommands(l, c, ssh, ifce)
|
attachCommands(l, c, ssh, ifce)
|
||||||
|
|
||||||
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
|
|
||||||
|
|
||||||
return &Control{
|
return &Control{
|
||||||
state: StateReady,
|
state: StateReady,
|
||||||
f: ifce,
|
f: ifce,
|
||||||
@@ -280,7 +269,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
statsStart: stats.Start,
|
statsStart: stats.Start,
|
||||||
dnsStart: ds.Start,
|
dnsStart: ds.Start,
|
||||||
lighthouseStart: lightHouse.StartUpdateWorker,
|
lighthouseStart: lightHouse.StartUpdateWorker,
|
||||||
networkChangeStart: networkChanges.Start,
|
|
||||||
connectionManagerStart: connManager.Start,
|
connectionManagerStart: connManager.Start,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -13,8 +13,6 @@ type MessageMetrics struct {
|
|||||||
|
|
||||||
rxUnknown metrics.Counter
|
rxUnknown metrics.Counter
|
||||||
txUnknown metrics.Counter
|
txUnknown metrics.Counter
|
||||||
|
|
||||||
rxInvalid metrics.Counter
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *MessageMetrics) Rx(t header.MessageType, s header.MessageSubType, i int64) {
|
func (m *MessageMetrics) Rx(t header.MessageType, s header.MessageSubType, i int64) {
|
||||||
@@ -35,11 +33,6 @@ func (m *MessageMetrics) Tx(t header.MessageType, s header.MessageSubType, i int
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
func (m *MessageMetrics) RxInvalid(i int64) {
|
|
||||||
if m != nil && m.rxInvalid != nil {
|
|
||||||
m.rxInvalid.Inc(i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newMessageMetrics() *MessageMetrics {
|
func newMessageMetrics() *MessageMetrics {
|
||||||
gen := func(t string) [][]metrics.Counter {
|
gen := func(t string) [][]metrics.Counter {
|
||||||
@@ -63,7 +56,6 @@ func newMessageMetrics() *MessageMetrics {
|
|||||||
|
|
||||||
rxUnknown: metrics.GetOrRegisterCounter("messages.rx.other", nil),
|
rxUnknown: metrics.GetOrRegisterCounter("messages.rx.other", nil),
|
||||||
txUnknown: metrics.GetOrRegisterCounter("messages.tx.other", nil),
|
txUnknown: metrics.GetOrRegisterCounter("messages.tx.other", nil),
|
||||||
rxInvalid: metrics.GetOrRegisterCounter("messages.rx.invalid", nil),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+632
-45
@@ -124,7 +124,7 @@ func (x NebulaControl_MessageType) String() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (NebulaControl_MessageType) EnumDescriptor() ([]byte, []int) {
|
func (NebulaControl_MessageType) EnumDescriptor() ([]byte, []int) {
|
||||||
return fileDescriptor_2d65afa7693df5ef, []int{6, 0}
|
return fileDescriptor_2d65afa7693df5ef, []int{8, 0}
|
||||||
}
|
}
|
||||||
|
|
||||||
type NebulaMeta struct {
|
type NebulaMeta struct {
|
||||||
@@ -489,6 +489,142 @@ func (m *NebulaPing) GetTime() uint64 {
|
|||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type NebulaHandshake struct {
|
||||||
|
Details *NebulaHandshakeDetails `protobuf:"bytes,1,opt,name=Details,proto3" json:"Details,omitempty"`
|
||||||
|
Hmac []byte `protobuf:"bytes,2,opt,name=Hmac,proto3" json:"Hmac,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshake) Reset() { *m = NebulaHandshake{} }
|
||||||
|
func (m *NebulaHandshake) String() string { return proto.CompactTextString(m) }
|
||||||
|
func (*NebulaHandshake) ProtoMessage() {}
|
||||||
|
func (*NebulaHandshake) Descriptor() ([]byte, []int) {
|
||||||
|
return fileDescriptor_2d65afa7693df5ef, []int{6}
|
||||||
|
}
|
||||||
|
func (m *NebulaHandshake) XXX_Unmarshal(b []byte) error {
|
||||||
|
return m.Unmarshal(b)
|
||||||
|
}
|
||||||
|
func (m *NebulaHandshake) XXX_Marshal(b []byte, deterministic bool) ([]byte, error) {
|
||||||
|
if deterministic {
|
||||||
|
return xxx_messageInfo_NebulaHandshake.Marshal(b, m, deterministic)
|
||||||
|
} else {
|
||||||
|
b = b[:cap(b)]
|
||||||
|
n, err := m.MarshalToSizedBuffer(b)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return b[:n], nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
func (m *NebulaHandshake) XXX_Merge(src proto.Message) {
|
||||||
|
xxx_messageInfo_NebulaHandshake.Merge(m, src)
|
||||||
|
}
|
||||||
|
func (m *NebulaHandshake) XXX_Size() int {
|
||||||
|
return m.Size()
|
||||||
|
}
|
||||||
|
func (m *NebulaHandshake) XXX_DiscardUnknown() {
|
||||||
|
xxx_messageInfo_NebulaHandshake.DiscardUnknown(m)
|
||||||
|
}
|
||||||
|
|
||||||
|
var xxx_messageInfo_NebulaHandshake proto.InternalMessageInfo
|
||||||
|
|
||||||
|
func (m *NebulaHandshake) GetDetails() *NebulaHandshakeDetails {
|
||||||
|
if m != nil {
|
||||||
|
return m.Details
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshake) GetHmac() []byte {
|
||||||
|
if m != nil {
|
||||||
|
return m.Hmac
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type NebulaHandshakeDetails struct {
|
||||||
|
Cert []byte `protobuf:"bytes,1,opt,name=Cert,proto3" json:"Cert,omitempty"`
|
||||||
|
InitiatorIndex uint32 `protobuf:"varint,2,opt,name=InitiatorIndex,proto3" json:"InitiatorIndex,omitempty"`
|
||||||
|
ResponderIndex uint32 `protobuf:"varint,3,opt,name=ResponderIndex,proto3" json:"ResponderIndex,omitempty"`
|
||||||
|
Cookie uint64 `protobuf:"varint,4,opt,name=Cookie,proto3" json:"Cookie,omitempty"`
|
||||||
|
Time uint64 `protobuf:"varint,5,opt,name=Time,proto3" json:"Time,omitempty"`
|
||||||
|
CertVersion uint32 `protobuf:"varint,8,opt,name=CertVersion,proto3" json:"CertVersion,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshakeDetails) Reset() { *m = NebulaHandshakeDetails{} }
|
||||||
|
func (m *NebulaHandshakeDetails) String() string { return proto.CompactTextString(m) }
|
||||||
|
func (*NebulaHandshakeDetails) ProtoMessage() {}
|
||||||
|
func (*NebulaHandshakeDetails) Descriptor() ([]byte, []int) {
|
||||||
|
return fileDescriptor_2d65afa7693df5ef, []int{7}
|
||||||
|
}
|
||||||
|
func (m *NebulaHandshakeDetails) XXX_Unmarshal(b []byte) error {
|
||||||
|
return m.Unmarshal(b)
|
||||||
|
}
|
||||||
|
func (m *NebulaHandshakeDetails) XXX_Marshal(b []byte, deterministic bool) ([]byte, error) {
|
||||||
|
if deterministic {
|
||||||
|
return xxx_messageInfo_NebulaHandshakeDetails.Marshal(b, m, deterministic)
|
||||||
|
} else {
|
||||||
|
b = b[:cap(b)]
|
||||||
|
n, err := m.MarshalToSizedBuffer(b)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return b[:n], nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
func (m *NebulaHandshakeDetails) XXX_Merge(src proto.Message) {
|
||||||
|
xxx_messageInfo_NebulaHandshakeDetails.Merge(m, src)
|
||||||
|
}
|
||||||
|
func (m *NebulaHandshakeDetails) XXX_Size() int {
|
||||||
|
return m.Size()
|
||||||
|
}
|
||||||
|
func (m *NebulaHandshakeDetails) XXX_DiscardUnknown() {
|
||||||
|
xxx_messageInfo_NebulaHandshakeDetails.DiscardUnknown(m)
|
||||||
|
}
|
||||||
|
|
||||||
|
var xxx_messageInfo_NebulaHandshakeDetails proto.InternalMessageInfo
|
||||||
|
|
||||||
|
func (m *NebulaHandshakeDetails) GetCert() []byte {
|
||||||
|
if m != nil {
|
||||||
|
return m.Cert
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshakeDetails) GetInitiatorIndex() uint32 {
|
||||||
|
if m != nil {
|
||||||
|
return m.InitiatorIndex
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshakeDetails) GetResponderIndex() uint32 {
|
||||||
|
if m != nil {
|
||||||
|
return m.ResponderIndex
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshakeDetails) GetCookie() uint64 {
|
||||||
|
if m != nil {
|
||||||
|
return m.Cookie
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshakeDetails) GetTime() uint64 {
|
||||||
|
if m != nil {
|
||||||
|
return m.Time
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshakeDetails) GetCertVersion() uint32 {
|
||||||
|
if m != nil {
|
||||||
|
return m.CertVersion
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
type NebulaControl struct {
|
type NebulaControl struct {
|
||||||
Type NebulaControl_MessageType `protobuf:"varint,1,opt,name=Type,proto3,enum=nebula.NebulaControl_MessageType" json:"Type,omitempty"`
|
Type NebulaControl_MessageType `protobuf:"varint,1,opt,name=Type,proto3,enum=nebula.NebulaControl_MessageType" json:"Type,omitempty"`
|
||||||
InitiatorRelayIndex uint32 `protobuf:"varint,2,opt,name=InitiatorRelayIndex,proto3" json:"InitiatorRelayIndex,omitempty"`
|
InitiatorRelayIndex uint32 `protobuf:"varint,2,opt,name=InitiatorRelayIndex,proto3" json:"InitiatorRelayIndex,omitempty"`
|
||||||
@@ -503,7 +639,7 @@ func (m *NebulaControl) Reset() { *m = NebulaControl{} }
|
|||||||
func (m *NebulaControl) String() string { return proto.CompactTextString(m) }
|
func (m *NebulaControl) String() string { return proto.CompactTextString(m) }
|
||||||
func (*NebulaControl) ProtoMessage() {}
|
func (*NebulaControl) ProtoMessage() {}
|
||||||
func (*NebulaControl) Descriptor() ([]byte, []int) {
|
func (*NebulaControl) Descriptor() ([]byte, []int) {
|
||||||
return fileDescriptor_2d65afa7693df5ef, []int{6}
|
return fileDescriptor_2d65afa7693df5ef, []int{8}
|
||||||
}
|
}
|
||||||
func (m *NebulaControl) XXX_Unmarshal(b []byte) error {
|
func (m *NebulaControl) XXX_Unmarshal(b []byte) error {
|
||||||
return m.Unmarshal(b)
|
return m.Unmarshal(b)
|
||||||
@@ -593,55 +729,65 @@ func init() {
|
|||||||
proto.RegisterType((*V4AddrPort)(nil), "nebula.V4AddrPort")
|
proto.RegisterType((*V4AddrPort)(nil), "nebula.V4AddrPort")
|
||||||
proto.RegisterType((*V6AddrPort)(nil), "nebula.V6AddrPort")
|
proto.RegisterType((*V6AddrPort)(nil), "nebula.V6AddrPort")
|
||||||
proto.RegisterType((*NebulaPing)(nil), "nebula.NebulaPing")
|
proto.RegisterType((*NebulaPing)(nil), "nebula.NebulaPing")
|
||||||
|
proto.RegisterType((*NebulaHandshake)(nil), "nebula.NebulaHandshake")
|
||||||
|
proto.RegisterType((*NebulaHandshakeDetails)(nil), "nebula.NebulaHandshakeDetails")
|
||||||
proto.RegisterType((*NebulaControl)(nil), "nebula.NebulaControl")
|
proto.RegisterType((*NebulaControl)(nil), "nebula.NebulaControl")
|
||||||
}
|
}
|
||||||
|
|
||||||
func init() { proto.RegisterFile("nebula.proto", fileDescriptor_2d65afa7693df5ef) }
|
func init() { proto.RegisterFile("nebula.proto", fileDescriptor_2d65afa7693df5ef) }
|
||||||
|
|
||||||
var fileDescriptor_2d65afa7693df5ef = []byte{
|
var fileDescriptor_2d65afa7693df5ef = []byte{
|
||||||
// 665 bytes of a gzipped FileDescriptorProto
|
// 785 bytes of a gzipped FileDescriptorProto
|
||||||
0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0xff, 0x84, 0x54, 0xcd, 0x6e, 0xd3, 0x5c,
|
0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0xff, 0x84, 0x55, 0xcd, 0x6e, 0xeb, 0x44,
|
||||||
0x10, 0x8d, 0x1d, 0x27, 0x69, 0x27, 0x4d, 0x3e, 0x7f, 0x53, 0x51, 0x12, 0x24, 0xac, 0xe0, 0x45,
|
0x14, 0x8e, 0x1d, 0x27, 0x4e, 0x4f, 0x7e, 0xae, 0x39, 0x15, 0xc1, 0x41, 0x22, 0x0a, 0x5e, 0x54,
|
||||||
0x55, 0xb1, 0x48, 0x51, 0x5a, 0xba, 0xa6, 0x2d, 0x42, 0xa9, 0xd4, 0x9f, 0x70, 0x55, 0x8a, 0xc4,
|
0x57, 0x2c, 0x72, 0x51, 0x5a, 0xae, 0x58, 0x72, 0x1b, 0x84, 0xd2, 0xaa, 0x3f, 0x61, 0x54, 0x8a,
|
||||||
0xce, 0xb5, 0x2f, 0x8d, 0x55, 0xc7, 0x37, 0xb5, 0x6f, 0x50, 0xf3, 0x16, 0x3c, 0x0c, 0x0f, 0x01,
|
0xc4, 0x06, 0xb9, 0xf6, 0xd0, 0x58, 0x71, 0x3c, 0xa9, 0x3d, 0x41, 0xcd, 0x5b, 0xf0, 0x30, 0x3c,
|
||||||
0xbb, 0x2e, 0x59, 0xa2, 0x66, 0xc9, 0x92, 0x17, 0x40, 0xf7, 0xfa, 0xbf, 0x31, 0xb0, 0xbb, 0x33,
|
0x04, 0xec, 0xba, 0x42, 0x2c, 0x51, 0xbb, 0x64, 0xc9, 0x0b, 0xa0, 0x19, 0xff, 0x27, 0x86, 0xbb,
|
||||||
0xe7, 0x9c, 0x99, 0xc9, 0xc9, 0x8c, 0x61, 0xcd, 0xa7, 0x97, 0x33, 0xcf, 0xea, 0x4f, 0x03, 0xc6,
|
0x9b, 0x73, 0xbe, 0xef, 0x3b, 0x73, 0xe6, 0xf3, 0x9c, 0x31, 0x74, 0x02, 0x7a, 0xb7, 0xf1, 0xed,
|
||||||
0x19, 0xd6, 0xa3, 0xc8, 0xfc, 0xa9, 0x02, 0x9c, 0xca, 0xe7, 0x09, 0xe5, 0x16, 0x0e, 0x40, 0x3b,
|
0xf1, 0x3a, 0x64, 0x9c, 0x61, 0x33, 0x8e, 0xac, 0xbf, 0x55, 0x80, 0x2b, 0xb9, 0xbc, 0xa4, 0xdc,
|
||||||
0x9f, 0x4f, 0x69, 0x47, 0xe9, 0x29, 0x5b, 0xed, 0x81, 0xd1, 0x8f, 0x35, 0x19, 0xa3, 0x7f, 0x42,
|
0xc6, 0x09, 0x68, 0x37, 0xdb, 0x35, 0x35, 0x95, 0x91, 0xf2, 0xba, 0x37, 0x19, 0x8e, 0x13, 0x4d,
|
||||||
0xc3, 0xd0, 0xba, 0xa2, 0x82, 0x45, 0x24, 0x17, 0x77, 0xa0, 0xf1, 0x9a, 0x72, 0xcb, 0xf5, 0xc2,
|
0xce, 0x18, 0x5f, 0xd2, 0x28, 0xb2, 0xef, 0xa9, 0x60, 0x11, 0xc9, 0xc5, 0x63, 0xd0, 0xbf, 0xa6,
|
||||||
0x8e, 0xda, 0x53, 0xb6, 0x9a, 0x83, 0xee, 0xb2, 0x2c, 0x26, 0x90, 0x84, 0x69, 0xfe, 0x52, 0xa0,
|
0xdc, 0xf6, 0xfc, 0xc8, 0x54, 0x47, 0xca, 0xeb, 0xf6, 0x64, 0xb0, 0x2f, 0x4b, 0x08, 0x24, 0x65,
|
||||||
0x99, 0x2b, 0x85, 0x2b, 0xa0, 0x9d, 0x32, 0x9f, 0xea, 0x15, 0x6c, 0xc1, 0xea, 0x90, 0x85, 0xfc,
|
0x5a, 0xff, 0x28, 0xd0, 0x2e, 0x94, 0xc2, 0x16, 0x68, 0x57, 0x2c, 0xa0, 0x46, 0x0d, 0xbb, 0x70,
|
||||||
0xed, 0x8c, 0x06, 0x73, 0x5d, 0x41, 0x84, 0x76, 0x1a, 0x12, 0x3a, 0xf5, 0xe6, 0xba, 0x8a, 0x4f,
|
0x30, 0x63, 0x11, 0xff, 0x76, 0x43, 0xc3, 0xad, 0xa1, 0x20, 0x42, 0x2f, 0x0b, 0x09, 0x5d, 0xfb,
|
||||||
0x60, 0x43, 0xe4, 0xde, 0x4d, 0x1d, 0x8b, 0xd3, 0x53, 0xc6, 0xdd, 0x8f, 0xae, 0x6d, 0x71, 0x97,
|
0x5b, 0x43, 0xc5, 0x8f, 0xa1, 0x2f, 0x72, 0xdf, 0xad, 0x5d, 0x9b, 0xd3, 0x2b, 0xc6, 0xbd, 0x9f,
|
||||||
0xf9, 0x7a, 0x15, 0xbb, 0xf0, 0x48, 0x60, 0x27, 0xec, 0x13, 0x75, 0x0a, 0x90, 0x96, 0x40, 0xa3,
|
0x3c, 0xc7, 0xe6, 0x1e, 0x0b, 0x8c, 0x3a, 0x0e, 0xe0, 0x43, 0x81, 0x5d, 0xb2, 0x9f, 0xa9, 0x5b,
|
||||||
0x99, 0x6f, 0x8f, 0x0b, 0x50, 0x0d, 0xdb, 0x00, 0x02, 0x7a, 0x3f, 0x66, 0xd6, 0xc4, 0xd5, 0xeb,
|
0x82, 0xb4, 0x14, 0x9a, 0x6f, 0x02, 0x67, 0x51, 0x82, 0x1a, 0xd8, 0x03, 0x10, 0xd0, 0xf7, 0x0b,
|
||||||
0xb8, 0x0e, 0xff, 0x65, 0x71, 0xd4, 0xb6, 0x21, 0x26, 0x1b, 0x59, 0x7c, 0x7c, 0x38, 0xa6, 0xf6,
|
0x66, 0xaf, 0x3c, 0xa3, 0x89, 0x87, 0xf0, 0x2a, 0x8f, 0xe3, 0x6d, 0x75, 0xd1, 0xd9, 0xdc, 0xe6,
|
||||||
0xb5, 0xbe, 0x22, 0x26, 0x4b, 0xc3, 0x88, 0xb2, 0x8a, 0x4f, 0xa1, 0x5b, 0x3e, 0xd9, 0xbe, 0x7d,
|
0x8b, 0xe9, 0x82, 0x3a, 0x4b, 0xa3, 0x25, 0x3a, 0xcb, 0xc2, 0x98, 0x72, 0x80, 0x9f, 0xc0, 0xa0,
|
||||||
0xad, 0x83, 0xf9, 0x4d, 0x85, 0xff, 0x97, 0x4c, 0x41, 0x13, 0xe0, 0xcc, 0x73, 0x2e, 0xa6, 0xfe,
|
0xba, 0xb3, 0x77, 0xce, 0xd2, 0x00, 0xeb, 0x77, 0x15, 0x3e, 0xd8, 0x33, 0x05, 0x2d, 0x80, 0x6b,
|
||||||
0xbe, 0xe3, 0x04, 0xd2, 0xfa, 0xd6, 0x81, 0xda, 0x51, 0x48, 0x2e, 0x8b, 0x9b, 0xd0, 0x48, 0x08,
|
0xdf, 0xbd, 0x5d, 0x07, 0xef, 0x5c, 0x37, 0x94, 0xd6, 0x77, 0x4f, 0x55, 0x53, 0x21, 0x85, 0x2c,
|
||||||
0x75, 0x69, 0xf2, 0x5a, 0x62, 0xb2, 0xc8, 0x91, 0x04, 0xc4, 0x3e, 0xe8, 0x67, 0x9e, 0x43, 0xa8,
|
0x1e, 0x81, 0x9e, 0x12, 0x9a, 0xd2, 0xe4, 0x4e, 0x6a, 0xb2, 0xc8, 0x91, 0x14, 0xc4, 0x31, 0x18,
|
||||||
0x67, 0xcd, 0xe3, 0x54, 0xd8, 0xa9, 0xf5, 0xaa, 0x71, 0xc5, 0x25, 0x0c, 0x07, 0xd0, 0x2a, 0x92,
|
0xd7, 0xbe, 0x4b, 0xa8, 0x6f, 0x6f, 0x93, 0x54, 0x64, 0x36, 0x46, 0xf5, 0xa4, 0xe2, 0x1e, 0x86,
|
||||||
0x1b, 0xbd, 0xea, 0x52, 0xf5, 0x22, 0x05, 0x77, 0xa1, 0x79, 0xb1, 0x2b, 0x9e, 0x23, 0x16, 0x70,
|
0x13, 0xe8, 0x96, 0xc9, 0xfa, 0xa8, 0xbe, 0x57, 0xbd, 0x4c, 0xc1, 0x13, 0x68, 0xdf, 0x9e, 0x88,
|
||||||
0xf1, 0xa7, 0x0b, 0x05, 0x26, 0x8a, 0x0c, 0x22, 0x79, 0x9a, 0x54, 0xed, 0x65, 0x2a, 0xed, 0x81,
|
0xe5, 0x9c, 0x85, 0x5c, 0x7c, 0x74, 0xa1, 0xc0, 0x54, 0x91, 0x43, 0xa4, 0x48, 0x93, 0xaa, 0xb7,
|
||||||
0x6a, 0x2f, 0xa7, 0xca, 0x68, 0xd8, 0x81, 0x86, 0xcd, 0x66, 0x3e, 0xa7, 0x41, 0xa7, 0x2a, 0x8c,
|
0xb9, 0x4a, 0xdb, 0x51, 0xbd, 0x2d, 0xa8, 0x72, 0x1a, 0x9a, 0xa0, 0x3b, 0x6c, 0x13, 0x70, 0x1a,
|
||||||
0x21, 0x49, 0x68, 0x6e, 0x82, 0x26, 0x7f, 0x71, 0x1b, 0xd4, 0xa1, 0x2b, 0x5d, 0xd3, 0x88, 0x3a,
|
0x9a, 0x75, 0x61, 0x0c, 0x49, 0x43, 0xeb, 0x08, 0x34, 0x79, 0xe2, 0x1e, 0xa8, 0x33, 0x4f, 0xba,
|
||||||
0x74, 0x45, 0x7c, 0xcc, 0xe4, 0x26, 0x6a, 0x44, 0x3d, 0x66, 0xe6, 0x2e, 0x40, 0x36, 0x06, 0x62,
|
0xa6, 0x11, 0x75, 0xe6, 0x89, 0xf8, 0x82, 0xc9, 0x9b, 0xa8, 0x11, 0xf5, 0x82, 0x59, 0x27, 0x00,
|
||||||
0xa4, 0x8a, 0x5c, 0x26, 0x51, 0x05, 0x04, 0x4d, 0x60, 0x52, 0xd3, 0x22, 0xf2, 0x6d, 0xbe, 0x02,
|
0x79, 0x1b, 0x88, 0xb1, 0x2a, 0x76, 0x99, 0xc4, 0x15, 0x10, 0x34, 0x81, 0x49, 0x4d, 0x97, 0xc8,
|
||||||
0xc8, 0xc6, 0xf8, 0x57, 0x8f, 0xb4, 0x42, 0x35, 0x57, 0xe1, 0x36, 0x39, 0xac, 0x91, 0xeb, 0x5f,
|
0xb5, 0xf5, 0x15, 0x40, 0xde, 0xc6, 0xfb, 0xf6, 0xc8, 0x2a, 0xd4, 0x0b, 0x15, 0x1e, 0xd3, 0xc1,
|
||||||
0xfd, 0xfd, 0xb0, 0x04, 0xa3, 0xe4, 0xb0, 0x10, 0xb4, 0x73, 0x77, 0x42, 0xe3, 0x3e, 0xf2, 0x6d,
|
0x9a, 0x7b, 0xc1, 0xfd, 0xff, 0x0f, 0x96, 0x60, 0x54, 0x0c, 0x16, 0x82, 0x76, 0xe3, 0xad, 0x68,
|
||||||
0x9a, 0x4b, 0x67, 0x23, 0xc4, 0x7a, 0x05, 0x57, 0xa1, 0x16, 0x2d, 0xa1, 0x62, 0x7e, 0xa9, 0x42,
|
0xb2, 0x8f, 0x5c, 0x5b, 0xd6, 0xde, 0xd8, 0x08, 0xb1, 0x51, 0xc3, 0x03, 0x68, 0xc4, 0x97, 0x50,
|
||||||
0x2b, 0x2a, 0x7c, 0xc8, 0x7c, 0x1e, 0x30, 0x0f, 0x5f, 0x16, 0xba, 0x3f, 0x2b, 0x76, 0x8f, 0x49,
|
0xb1, 0x7e, 0x84, 0x57, 0x71, 0xdd, 0x99, 0x1d, 0xb8, 0xd1, 0xc2, 0x5e, 0x52, 0xfc, 0x32, 0x9f,
|
||||||
0x25, 0x03, 0xbc, 0x80, 0xf5, 0x23, 0xdf, 0xe5, 0xae, 0xc5, 0x59, 0x20, 0x57, 0xe0, 0xc8, 0x77,
|
0x51, 0x45, 0x5e, 0x9f, 0x9d, 0x0e, 0x32, 0xe6, 0xee, 0xa0, 0x8a, 0x26, 0x66, 0x2b, 0xdb, 0x91,
|
||||||
0xe8, 0x6d, 0xec, 0x53, 0x19, 0x24, 0x14, 0x84, 0x86, 0x53, 0xe6, 0x3b, 0x34, 0xaf, 0x88, 0x7c,
|
0x4d, 0x74, 0x88, 0x5c, 0x5b, 0x7f, 0x28, 0xd0, 0xaf, 0xd6, 0x09, 0xfa, 0x94, 0x86, 0x5c, 0xee,
|
||||||
0x29, 0x83, 0xf0, 0x39, 0xb4, 0x93, 0xa5, 0x3c, 0x67, 0xf2, 0xaf, 0xd1, 0xd2, 0x03, 0x78, 0x80,
|
0xd2, 0x21, 0x72, 0x8d, 0x47, 0xd0, 0x3b, 0x0b, 0x3c, 0xee, 0xd9, 0x9c, 0x85, 0x67, 0x81, 0x4b,
|
||||||
0xe4, 0x97, 0xfb, 0x4d, 0xc0, 0x26, 0x92, 0x5d, 0x4b, 0xd9, 0x4b, 0x18, 0xf6, 0xa1, 0x99, 0x2f,
|
0x1f, 0x13, 0xa7, 0x77, 0xb2, 0x82, 0x47, 0x68, 0xb4, 0x66, 0x81, 0x4b, 0x13, 0x5e, 0xec, 0xe7,
|
||||||
0x5c, 0x76, 0x38, 0x79, 0x42, 0x7a, 0x0c, 0x69, 0xf1, 0x46, 0x89, 0xa2, 0x48, 0x31, 0x87, 0x7f,
|
0x4e, 0x16, 0xfb, 0xd0, 0x9c, 0x32, 0xb6, 0xf4, 0xa8, 0xa9, 0x49, 0x67, 0x92, 0x28, 0xf3, 0xab,
|
||||||
0xfa, 0x8e, 0x6d, 0x00, 0x1e, 0x06, 0xd4, 0xe2, 0x54, 0xf2, 0x09, 0xbd, 0x99, 0xd1, 0x90, 0xeb,
|
0x91, 0xfb, 0x85, 0x23, 0x68, 0x8b, 0x1e, 0x6e, 0x69, 0x18, 0x79, 0x2c, 0x30, 0x5b, 0xb2, 0x60,
|
||||||
0x0a, 0x3e, 0x86, 0xf5, 0x42, 0x5e, 0x58, 0x12, 0x52, 0x5d, 0x3d, 0xd8, 0xf9, 0x7a, 0x6f, 0x28,
|
0x31, 0x75, 0xae, 0xb5, 0x9a, 0x86, 0x7e, 0xae, 0xb5, 0x74, 0xa3, 0x65, 0xfd, 0x5a, 0x87, 0x6e,
|
||||||
0x77, 0xf7, 0x86, 0xf2, 0xe3, 0xde, 0x50, 0x3e, 0x2f, 0x8c, 0xca, 0xdd, 0xc2, 0xa8, 0x7c, 0x5f,
|
0x7c, 0xb0, 0x29, 0x0b, 0x78, 0xc8, 0x7c, 0xfc, 0xa2, 0xf4, 0xdd, 0x3e, 0x2d, 0xbb, 0x96, 0x90,
|
||||||
0x18, 0x95, 0x0f, 0xdd, 0x2b, 0x97, 0x8f, 0x67, 0x97, 0x7d, 0x9b, 0x4d, 0xb6, 0x43, 0xcf, 0xb2,
|
0x2a, 0x3e, 0xdd, 0xe7, 0x70, 0x98, 0x1d, 0x4e, 0x0e, 0x4f, 0xf1, 0xdc, 0x55, 0x90, 0x50, 0x64,
|
||||||
0xaf, 0xc7, 0x37, 0xdb, 0xd1, 0x48, 0x97, 0x75, 0xf9, 0x39, 0xdf, 0xf9, 0x1d, 0x00, 0x00, 0xff,
|
0xc7, 0x2c, 0x28, 0x62, 0x07, 0xaa, 0x20, 0xfc, 0x0c, 0x7a, 0xe9, 0x38, 0xdf, 0x30, 0x79, 0xa9,
|
||||||
0xff, 0x51, 0x0a, 0xe3, 0xd7, 0xde, 0x05, 0x00, 0x00,
|
0xb5, 0xec, 0xe9, 0xd8, 0x41, 0x8a, 0xcf, 0xc2, 0x37, 0x21, 0x5b, 0x49, 0x76, 0x23, 0x63, 0xef,
|
||||||
|
0x61, 0x38, 0x86, 0x76, 0xb1, 0x70, 0xd5, 0x93, 0x53, 0x24, 0x64, 0xcf, 0x48, 0x56, 0x5c, 0xaf,
|
||||||
|
0x50, 0x94, 0x29, 0xd6, 0xec, 0xbf, 0xfe, 0x00, 0x7d, 0xc0, 0x69, 0x48, 0x6d, 0x4e, 0x25, 0x9f,
|
||||||
|
0xd0, 0x87, 0x0d, 0x8d, 0xb8, 0xa1, 0xe0, 0x47, 0x70, 0x58, 0xca, 0x0b, 0x4b, 0x22, 0x6a, 0xa8,
|
||||||
|
0xa7, 0xc7, 0xbf, 0x3d, 0x0f, 0x95, 0xa7, 0xe7, 0xa1, 0xf2, 0xd7, 0xf3, 0x50, 0xf9, 0xe5, 0x65,
|
||||||
|
0x58, 0x7b, 0x7a, 0x19, 0xd6, 0xfe, 0x7c, 0x19, 0xd6, 0x7e, 0x18, 0xdc, 0x7b, 0x7c, 0xb1, 0xb9,
|
||||||
|
0x1b, 0x3b, 0x6c, 0xf5, 0x26, 0xf2, 0x6d, 0x67, 0xb9, 0x78, 0x78, 0x13, 0xb7, 0x74, 0xd7, 0x94,
|
||||||
|
0x3f, 0xc2, 0xe3, 0x7f, 0x03, 0x00, 0x00, 0xff, 0xff, 0xea, 0x6f, 0xbc, 0x50, 0x18, 0x07, 0x00,
|
||||||
|
0x00,
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *NebulaMeta) Marshal() (dAtA []byte, err error) {
|
func (m *NebulaMeta) Marshal() (dAtA []byte, err error) {
|
||||||
@@ -926,6 +1072,103 @@ func (m *NebulaPing) MarshalToSizedBuffer(dAtA []byte) (int, error) {
|
|||||||
return len(dAtA) - i, nil
|
return len(dAtA) - i, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshake) Marshal() (dAtA []byte, err error) {
|
||||||
|
size := m.Size()
|
||||||
|
dAtA = make([]byte, size)
|
||||||
|
n, err := m.MarshalToSizedBuffer(dAtA[:size])
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return dAtA[:n], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshake) MarshalTo(dAtA []byte) (int, error) {
|
||||||
|
size := m.Size()
|
||||||
|
return m.MarshalToSizedBuffer(dAtA[:size])
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshake) MarshalToSizedBuffer(dAtA []byte) (int, error) {
|
||||||
|
i := len(dAtA)
|
||||||
|
_ = i
|
||||||
|
var l int
|
||||||
|
_ = l
|
||||||
|
if len(m.Hmac) > 0 {
|
||||||
|
i -= len(m.Hmac)
|
||||||
|
copy(dAtA[i:], m.Hmac)
|
||||||
|
i = encodeVarintNebula(dAtA, i, uint64(len(m.Hmac)))
|
||||||
|
i--
|
||||||
|
dAtA[i] = 0x12
|
||||||
|
}
|
||||||
|
if m.Details != nil {
|
||||||
|
{
|
||||||
|
size, err := m.Details.MarshalToSizedBuffer(dAtA[:i])
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
i -= size
|
||||||
|
i = encodeVarintNebula(dAtA, i, uint64(size))
|
||||||
|
}
|
||||||
|
i--
|
||||||
|
dAtA[i] = 0xa
|
||||||
|
}
|
||||||
|
return len(dAtA) - i, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshakeDetails) Marshal() (dAtA []byte, err error) {
|
||||||
|
size := m.Size()
|
||||||
|
dAtA = make([]byte, size)
|
||||||
|
n, err := m.MarshalToSizedBuffer(dAtA[:size])
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return dAtA[:n], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshakeDetails) MarshalTo(dAtA []byte) (int, error) {
|
||||||
|
size := m.Size()
|
||||||
|
return m.MarshalToSizedBuffer(dAtA[:size])
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshakeDetails) MarshalToSizedBuffer(dAtA []byte) (int, error) {
|
||||||
|
i := len(dAtA)
|
||||||
|
_ = i
|
||||||
|
var l int
|
||||||
|
_ = l
|
||||||
|
if m.CertVersion != 0 {
|
||||||
|
i = encodeVarintNebula(dAtA, i, uint64(m.CertVersion))
|
||||||
|
i--
|
||||||
|
dAtA[i] = 0x40
|
||||||
|
}
|
||||||
|
if m.Time != 0 {
|
||||||
|
i = encodeVarintNebula(dAtA, i, uint64(m.Time))
|
||||||
|
i--
|
||||||
|
dAtA[i] = 0x28
|
||||||
|
}
|
||||||
|
if m.Cookie != 0 {
|
||||||
|
i = encodeVarintNebula(dAtA, i, uint64(m.Cookie))
|
||||||
|
i--
|
||||||
|
dAtA[i] = 0x20
|
||||||
|
}
|
||||||
|
if m.ResponderIndex != 0 {
|
||||||
|
i = encodeVarintNebula(dAtA, i, uint64(m.ResponderIndex))
|
||||||
|
i--
|
||||||
|
dAtA[i] = 0x18
|
||||||
|
}
|
||||||
|
if m.InitiatorIndex != 0 {
|
||||||
|
i = encodeVarintNebula(dAtA, i, uint64(m.InitiatorIndex))
|
||||||
|
i--
|
||||||
|
dAtA[i] = 0x10
|
||||||
|
}
|
||||||
|
if len(m.Cert) > 0 {
|
||||||
|
i -= len(m.Cert)
|
||||||
|
copy(dAtA[i:], m.Cert)
|
||||||
|
i = encodeVarintNebula(dAtA, i, uint64(len(m.Cert)))
|
||||||
|
i--
|
||||||
|
dAtA[i] = 0xa
|
||||||
|
}
|
||||||
|
return len(dAtA) - i, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (m *NebulaControl) Marshal() (dAtA []byte, err error) {
|
func (m *NebulaControl) Marshal() (dAtA []byte, err error) {
|
||||||
size := m.Size()
|
size := m.Size()
|
||||||
dAtA = make([]byte, size)
|
dAtA = make([]byte, size)
|
||||||
@@ -1132,6 +1375,51 @@ func (m *NebulaPing) Size() (n int) {
|
|||||||
return n
|
return n
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshake) Size() (n int) {
|
||||||
|
if m == nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
var l int
|
||||||
|
_ = l
|
||||||
|
if m.Details != nil {
|
||||||
|
l = m.Details.Size()
|
||||||
|
n += 1 + l + sovNebula(uint64(l))
|
||||||
|
}
|
||||||
|
l = len(m.Hmac)
|
||||||
|
if l > 0 {
|
||||||
|
n += 1 + l + sovNebula(uint64(l))
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *NebulaHandshakeDetails) Size() (n int) {
|
||||||
|
if m == nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
var l int
|
||||||
|
_ = l
|
||||||
|
l = len(m.Cert)
|
||||||
|
if l > 0 {
|
||||||
|
n += 1 + l + sovNebula(uint64(l))
|
||||||
|
}
|
||||||
|
if m.InitiatorIndex != 0 {
|
||||||
|
n += 1 + sovNebula(uint64(m.InitiatorIndex))
|
||||||
|
}
|
||||||
|
if m.ResponderIndex != 0 {
|
||||||
|
n += 1 + sovNebula(uint64(m.ResponderIndex))
|
||||||
|
}
|
||||||
|
if m.Cookie != 0 {
|
||||||
|
n += 1 + sovNebula(uint64(m.Cookie))
|
||||||
|
}
|
||||||
|
if m.Time != 0 {
|
||||||
|
n += 1 + sovNebula(uint64(m.Time))
|
||||||
|
}
|
||||||
|
if m.CertVersion != 0 {
|
||||||
|
n += 1 + sovNebula(uint64(m.CertVersion))
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
func (m *NebulaControl) Size() (n int) {
|
func (m *NebulaControl) Size() (n int) {
|
||||||
if m == nil {
|
if m == nil {
|
||||||
return 0
|
return 0
|
||||||
@@ -1948,6 +2236,305 @@ func (m *NebulaPing) Unmarshal(dAtA []byte) error {
|
|||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
func (m *NebulaHandshake) Unmarshal(dAtA []byte) error {
|
||||||
|
l := len(dAtA)
|
||||||
|
iNdEx := 0
|
||||||
|
for iNdEx < l {
|
||||||
|
preIndex := iNdEx
|
||||||
|
var wire uint64
|
||||||
|
for shift := uint(0); ; shift += 7 {
|
||||||
|
if shift >= 64 {
|
||||||
|
return ErrIntOverflowNebula
|
||||||
|
}
|
||||||
|
if iNdEx >= l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
b := dAtA[iNdEx]
|
||||||
|
iNdEx++
|
||||||
|
wire |= uint64(b&0x7F) << shift
|
||||||
|
if b < 0x80 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
fieldNum := int32(wire >> 3)
|
||||||
|
wireType := int(wire & 0x7)
|
||||||
|
if wireType == 4 {
|
||||||
|
return fmt.Errorf("proto: NebulaHandshake: wiretype end group for non-group")
|
||||||
|
}
|
||||||
|
if fieldNum <= 0 {
|
||||||
|
return fmt.Errorf("proto: NebulaHandshake: illegal tag %d (wire type %d)", fieldNum, wire)
|
||||||
|
}
|
||||||
|
switch fieldNum {
|
||||||
|
case 1:
|
||||||
|
if wireType != 2 {
|
||||||
|
return fmt.Errorf("proto: wrong wireType = %d for field Details", wireType)
|
||||||
|
}
|
||||||
|
var msglen int
|
||||||
|
for shift := uint(0); ; shift += 7 {
|
||||||
|
if shift >= 64 {
|
||||||
|
return ErrIntOverflowNebula
|
||||||
|
}
|
||||||
|
if iNdEx >= l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
b := dAtA[iNdEx]
|
||||||
|
iNdEx++
|
||||||
|
msglen |= int(b&0x7F) << shift
|
||||||
|
if b < 0x80 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if msglen < 0 {
|
||||||
|
return ErrInvalidLengthNebula
|
||||||
|
}
|
||||||
|
postIndex := iNdEx + msglen
|
||||||
|
if postIndex < 0 {
|
||||||
|
return ErrInvalidLengthNebula
|
||||||
|
}
|
||||||
|
if postIndex > l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
if m.Details == nil {
|
||||||
|
m.Details = &NebulaHandshakeDetails{}
|
||||||
|
}
|
||||||
|
if err := m.Details.Unmarshal(dAtA[iNdEx:postIndex]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
iNdEx = postIndex
|
||||||
|
case 2:
|
||||||
|
if wireType != 2 {
|
||||||
|
return fmt.Errorf("proto: wrong wireType = %d for field Hmac", wireType)
|
||||||
|
}
|
||||||
|
var byteLen int
|
||||||
|
for shift := uint(0); ; shift += 7 {
|
||||||
|
if shift >= 64 {
|
||||||
|
return ErrIntOverflowNebula
|
||||||
|
}
|
||||||
|
if iNdEx >= l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
b := dAtA[iNdEx]
|
||||||
|
iNdEx++
|
||||||
|
byteLen |= int(b&0x7F) << shift
|
||||||
|
if b < 0x80 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if byteLen < 0 {
|
||||||
|
return ErrInvalidLengthNebula
|
||||||
|
}
|
||||||
|
postIndex := iNdEx + byteLen
|
||||||
|
if postIndex < 0 {
|
||||||
|
return ErrInvalidLengthNebula
|
||||||
|
}
|
||||||
|
if postIndex > l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
m.Hmac = append(m.Hmac[:0], dAtA[iNdEx:postIndex]...)
|
||||||
|
if m.Hmac == nil {
|
||||||
|
m.Hmac = []byte{}
|
||||||
|
}
|
||||||
|
iNdEx = postIndex
|
||||||
|
default:
|
||||||
|
iNdEx = preIndex
|
||||||
|
skippy, err := skipNebula(dAtA[iNdEx:])
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if (skippy < 0) || (iNdEx+skippy) < 0 {
|
||||||
|
return ErrInvalidLengthNebula
|
||||||
|
}
|
||||||
|
if (iNdEx + skippy) > l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
iNdEx += skippy
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if iNdEx > l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (m *NebulaHandshakeDetails) Unmarshal(dAtA []byte) error {
|
||||||
|
l := len(dAtA)
|
||||||
|
iNdEx := 0
|
||||||
|
for iNdEx < l {
|
||||||
|
preIndex := iNdEx
|
||||||
|
var wire uint64
|
||||||
|
for shift := uint(0); ; shift += 7 {
|
||||||
|
if shift >= 64 {
|
||||||
|
return ErrIntOverflowNebula
|
||||||
|
}
|
||||||
|
if iNdEx >= l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
b := dAtA[iNdEx]
|
||||||
|
iNdEx++
|
||||||
|
wire |= uint64(b&0x7F) << shift
|
||||||
|
if b < 0x80 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
fieldNum := int32(wire >> 3)
|
||||||
|
wireType := int(wire & 0x7)
|
||||||
|
if wireType == 4 {
|
||||||
|
return fmt.Errorf("proto: NebulaHandshakeDetails: wiretype end group for non-group")
|
||||||
|
}
|
||||||
|
if fieldNum <= 0 {
|
||||||
|
return fmt.Errorf("proto: NebulaHandshakeDetails: illegal tag %d (wire type %d)", fieldNum, wire)
|
||||||
|
}
|
||||||
|
switch fieldNum {
|
||||||
|
case 1:
|
||||||
|
if wireType != 2 {
|
||||||
|
return fmt.Errorf("proto: wrong wireType = %d for field Cert", wireType)
|
||||||
|
}
|
||||||
|
var byteLen int
|
||||||
|
for shift := uint(0); ; shift += 7 {
|
||||||
|
if shift >= 64 {
|
||||||
|
return ErrIntOverflowNebula
|
||||||
|
}
|
||||||
|
if iNdEx >= l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
b := dAtA[iNdEx]
|
||||||
|
iNdEx++
|
||||||
|
byteLen |= int(b&0x7F) << shift
|
||||||
|
if b < 0x80 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if byteLen < 0 {
|
||||||
|
return ErrInvalidLengthNebula
|
||||||
|
}
|
||||||
|
postIndex := iNdEx + byteLen
|
||||||
|
if postIndex < 0 {
|
||||||
|
return ErrInvalidLengthNebula
|
||||||
|
}
|
||||||
|
if postIndex > l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
m.Cert = append(m.Cert[:0], dAtA[iNdEx:postIndex]...)
|
||||||
|
if m.Cert == nil {
|
||||||
|
m.Cert = []byte{}
|
||||||
|
}
|
||||||
|
iNdEx = postIndex
|
||||||
|
case 2:
|
||||||
|
if wireType != 0 {
|
||||||
|
return fmt.Errorf("proto: wrong wireType = %d for field InitiatorIndex", wireType)
|
||||||
|
}
|
||||||
|
m.InitiatorIndex = 0
|
||||||
|
for shift := uint(0); ; shift += 7 {
|
||||||
|
if shift >= 64 {
|
||||||
|
return ErrIntOverflowNebula
|
||||||
|
}
|
||||||
|
if iNdEx >= l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
b := dAtA[iNdEx]
|
||||||
|
iNdEx++
|
||||||
|
m.InitiatorIndex |= uint32(b&0x7F) << shift
|
||||||
|
if b < 0x80 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case 3:
|
||||||
|
if wireType != 0 {
|
||||||
|
return fmt.Errorf("proto: wrong wireType = %d for field ResponderIndex", wireType)
|
||||||
|
}
|
||||||
|
m.ResponderIndex = 0
|
||||||
|
for shift := uint(0); ; shift += 7 {
|
||||||
|
if shift >= 64 {
|
||||||
|
return ErrIntOverflowNebula
|
||||||
|
}
|
||||||
|
if iNdEx >= l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
b := dAtA[iNdEx]
|
||||||
|
iNdEx++
|
||||||
|
m.ResponderIndex |= uint32(b&0x7F) << shift
|
||||||
|
if b < 0x80 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case 4:
|
||||||
|
if wireType != 0 {
|
||||||
|
return fmt.Errorf("proto: wrong wireType = %d for field Cookie", wireType)
|
||||||
|
}
|
||||||
|
m.Cookie = 0
|
||||||
|
for shift := uint(0); ; shift += 7 {
|
||||||
|
if shift >= 64 {
|
||||||
|
return ErrIntOverflowNebula
|
||||||
|
}
|
||||||
|
if iNdEx >= l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
b := dAtA[iNdEx]
|
||||||
|
iNdEx++
|
||||||
|
m.Cookie |= uint64(b&0x7F) << shift
|
||||||
|
if b < 0x80 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case 5:
|
||||||
|
if wireType != 0 {
|
||||||
|
return fmt.Errorf("proto: wrong wireType = %d for field Time", wireType)
|
||||||
|
}
|
||||||
|
m.Time = 0
|
||||||
|
for shift := uint(0); ; shift += 7 {
|
||||||
|
if shift >= 64 {
|
||||||
|
return ErrIntOverflowNebula
|
||||||
|
}
|
||||||
|
if iNdEx >= l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
b := dAtA[iNdEx]
|
||||||
|
iNdEx++
|
||||||
|
m.Time |= uint64(b&0x7F) << shift
|
||||||
|
if b < 0x80 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case 8:
|
||||||
|
if wireType != 0 {
|
||||||
|
return fmt.Errorf("proto: wrong wireType = %d for field CertVersion", wireType)
|
||||||
|
}
|
||||||
|
m.CertVersion = 0
|
||||||
|
for shift := uint(0); ; shift += 7 {
|
||||||
|
if shift >= 64 {
|
||||||
|
return ErrIntOverflowNebula
|
||||||
|
}
|
||||||
|
if iNdEx >= l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
b := dAtA[iNdEx]
|
||||||
|
iNdEx++
|
||||||
|
m.CertVersion |= uint32(b&0x7F) << shift
|
||||||
|
if b < 0x80 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
iNdEx = preIndex
|
||||||
|
skippy, err := skipNebula(dAtA[iNdEx:])
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if (skippy < 0) || (iNdEx+skippy) < 0 {
|
||||||
|
return ErrInvalidLengthNebula
|
||||||
|
}
|
||||||
|
if (iNdEx + skippy) > l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
iNdEx += skippy
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if iNdEx > l {
|
||||||
|
return io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
func (m *NebulaControl) Unmarshal(dAtA []byte) error {
|
func (m *NebulaControl) Unmarshal(dAtA []byte) error {
|
||||||
l := len(dAtA)
|
l := len(dAtA)
|
||||||
iNdEx := 0
|
iNdEx := 0
|
||||||
|
|||||||
+15
-3
@@ -60,9 +60,21 @@ message NebulaPing {
|
|||||||
uint64 Time = 2;
|
uint64 Time = 2;
|
||||||
}
|
}
|
||||||
|
|
||||||
// NebulaHandshake / NebulaHandshakeDetails moved to
|
message NebulaHandshake {
|
||||||
// handshake/handshake.proto. The handshake package speaks that wire format
|
NebulaHandshakeDetails Details = 1;
|
||||||
// directly via a hand-written encoder/decoder.
|
bytes Hmac = 2;
|
||||||
|
}
|
||||||
|
|
||||||
|
message NebulaHandshakeDetails {
|
||||||
|
bytes Cert = 1;
|
||||||
|
uint32 InitiatorIndex = 2;
|
||||||
|
uint32 ResponderIndex = 3;
|
||||||
|
uint64 Cookie = 4;
|
||||||
|
uint64 Time = 5;
|
||||||
|
uint32 CertVersion = 8;
|
||||||
|
// reserved for WIP multiport
|
||||||
|
reserved 6, 7;
|
||||||
|
}
|
||||||
|
|
||||||
message NebulaControl {
|
message NebulaControl {
|
||||||
enum MessageType {
|
enum MessageType {
|
||||||
|
|||||||
@@ -0,0 +1,73 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/cipher"
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
)
|
||||||
|
|
||||||
|
type endianness interface {
|
||||||
|
PutUint64(b []byte, v uint64)
|
||||||
|
}
|
||||||
|
|
||||||
|
var noiseEndianness endianness = binary.BigEndian
|
||||||
|
|
||||||
|
type NebulaCipherState struct {
|
||||||
|
c cipher.AEAD
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewNebulaCipherState(s *noise.CipherState) *NebulaCipherState {
|
||||||
|
x := s.Cipher()
|
||||||
|
return &NebulaCipherState{c: x.(cipher.AEAD)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// EncryptDanger encrypts and authenticates a given payload.
|
||||||
|
//
|
||||||
|
// out is a destination slice to hold the output of the EncryptDanger operation.
|
||||||
|
// - ad is additional data, which will be authenticated and appended to out, but not encrypted.
|
||||||
|
// - plaintext is encrypted, authenticated and appended to out.
|
||||||
|
// - n is a nonce value which must never be re-used with this key.
|
||||||
|
// - nb is a buffer used for temporary storage in the implementation of this call, which should
|
||||||
|
// be re-used by callers to minimize garbage collection.
|
||||||
|
func (s *NebulaCipherState) EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||||
|
if s != nil {
|
||||||
|
// TODO: Is this okay now that we have made messageCounter atomic?
|
||||||
|
// Alternative may be to split the counter space into ranges
|
||||||
|
//if n <= s.n {
|
||||||
|
// return nil, errors.New("CRITICAL: a duplicate counter value was used")
|
||||||
|
//}
|
||||||
|
//s.n = n
|
||||||
|
nb[0] = 0
|
||||||
|
nb[1] = 0
|
||||||
|
nb[2] = 0
|
||||||
|
nb[3] = 0
|
||||||
|
noiseEndianness.PutUint64(nb[4:], n)
|
||||||
|
out = s.c.Seal(out, nb, plaintext, ad)
|
||||||
|
//l.Debugf("Encryption: outlen: %d, nonce: %d, ad: %s, plainlen %d", len(out), n, ad, len(plaintext))
|
||||||
|
return out, nil
|
||||||
|
} else {
|
||||||
|
return nil, errors.New("no cipher state available to encrypt")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *NebulaCipherState) DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||||
|
if s != nil {
|
||||||
|
nb[0] = 0
|
||||||
|
nb[1] = 0
|
||||||
|
nb[2] = 0
|
||||||
|
nb[3] = 0
|
||||||
|
noiseEndianness.PutUint64(nb[4:], n)
|
||||||
|
return s.c.Open(out, nb, ciphertext, ad)
|
||||||
|
} else {
|
||||||
|
return []byte{}, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *NebulaCipherState) Overhead() int {
|
||||||
|
if s != nil {
|
||||||
|
return s.c.Overhead()
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
@@ -1,53 +0,0 @@
|
|||||||
package noiseutil
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/cipher"
|
|
||||||
"encoding/binary"
|
|
||||||
"errors"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
)
|
|
||||||
|
|
||||||
// CipherStateAESGCM is the data-plane wrapper for the AES-GCM AEAD cipher.
|
|
||||||
// AES-GCM uses big-endian nonce encoding per the Noise spec.
|
|
||||||
type CipherStateAESGCM struct {
|
|
||||||
c cipher.AEAD
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewCipherStateAESGCM extracts the underlying AEAD from the post-handshake noise.CipherState.
|
|
||||||
// The caller is responsible for ensuring the noise cipher is actually AES-GCM,
|
|
||||||
// otherwise the type assertion still succeeds but the nonce endianness will be wrong on the wire.
|
|
||||||
func NewCipherStateAESGCM(s *noise.CipherState) *CipherStateAESGCM {
|
|
||||||
return &CipherStateAESGCM{c: s.Cipher().(cipher.AEAD)}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *CipherStateAESGCM) EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error) {
|
|
||||||
if s == nil {
|
|
||||||
return nil, errors.New("no cipher state available to encrypt")
|
|
||||||
}
|
|
||||||
nb[0] = 0
|
|
||||||
nb[1] = 0
|
|
||||||
nb[2] = 0
|
|
||||||
nb[3] = 0
|
|
||||||
binary.BigEndian.PutUint64(nb[4:], n)
|
|
||||||
return s.c.Seal(out, nb, plaintext, ad), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *CipherStateAESGCM) DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error) {
|
|
||||||
if s == nil {
|
|
||||||
return []byte{}, nil
|
|
||||||
}
|
|
||||||
nb[0] = 0
|
|
||||||
nb[1] = 0
|
|
||||||
nb[2] = 0
|
|
||||||
nb[3] = 0
|
|
||||||
binary.BigEndian.PutUint64(nb[4:], n)
|
|
||||||
return s.c.Open(out, nb, ciphertext, ad)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *CipherStateAESGCM) Overhead() int {
|
|
||||||
if s == nil {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
return s.c.Overhead()
|
|
||||||
}
|
|
||||||
@@ -1,52 +0,0 @@
|
|||||||
package noiseutil
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/cipher"
|
|
||||||
"encoding/binary"
|
|
||||||
"errors"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
)
|
|
||||||
|
|
||||||
// CipherStateChaChaPoly is the data-plane wrapper for the ChaCha20-Poly1305 AEAD cipher.
|
|
||||||
// ChaCha20-Poly1305 uses little-endian nonce encoding per the Noise spec.
|
|
||||||
type CipherStateChaChaPoly struct {
|
|
||||||
c cipher.AEAD
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewCipherStateChaChaPoly extracts the underlying AEAD from the post-handshake noise.CipherState.
|
|
||||||
// The caller is responsible for ensuring the noise cipher is actually ChaCha20-Poly1305.
|
|
||||||
func NewCipherStateChaChaPoly(s *noise.CipherState) *CipherStateChaChaPoly {
|
|
||||||
return &CipherStateChaChaPoly{c: s.Cipher().(cipher.AEAD)}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *CipherStateChaChaPoly) EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error) {
|
|
||||||
if s == nil {
|
|
||||||
return nil, errors.New("no cipher state available to encrypt")
|
|
||||||
}
|
|
||||||
nb[0] = 0
|
|
||||||
nb[1] = 0
|
|
||||||
nb[2] = 0
|
|
||||||
nb[3] = 0
|
|
||||||
binary.LittleEndian.PutUint64(nb[4:], n)
|
|
||||||
return s.c.Seal(out, nb, plaintext, ad), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *CipherStateChaChaPoly) DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error) {
|
|
||||||
if s == nil {
|
|
||||||
return []byte{}, nil
|
|
||||||
}
|
|
||||||
nb[0] = 0
|
|
||||||
nb[1] = 0
|
|
||||||
nb[2] = 0
|
|
||||||
nb[3] = 0
|
|
||||||
binary.LittleEndian.PutUint64(nb[4:], n)
|
|
||||||
return s.c.Open(out, nb, ciphertext, ad)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *CipherStateChaChaPoly) Overhead() int {
|
|
||||||
if s == nil {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
return s.c.Overhead()
|
|
||||||
}
|
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user