mirror of
https://github.com/slackhq/nebula.git
synced 2026-09-30 07:36:37 +02:00
Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d6259aef68 | ||
|
|
7fc40557a8 | ||
|
|
bdaf1eb92f |
@@ -43,15 +43,8 @@ runs:
|
|||||||
with:
|
with:
|
||||||
role-to-assume: ${{ inputs.role }}
|
role-to-assume: ${{ inputs.role }}
|
||||||
aws-region: ${{ inputs.region }}
|
aws-region: ${{ inputs.region }}
|
||||||
# An STS secret key with special characters does not survive the
|
# Default is 12 retries to ride out IAM trust-policy propagation; once
|
||||||
# pwsh -> make -> MSYS sh -> aws.exe chain, and SigV4 then signs with a
|
# the role is stable we want a real misconfiguration to fail fast.
|
||||||
# key that no longer matches, so the first S3 upload fails with
|
|
||||||
# SignatureDoesNotMatch. Retries the assume until it comes back clean.
|
|
||||||
# Same fix as DefinedNet/dnclient#867.
|
|
||||||
special-characters-workaround: true
|
|
||||||
# Overridden by the workaround above and kept for whenever that goes:
|
|
||||||
# the default 12 rides out IAM trust-policy propagation, and once the
|
|
||||||
# role is stable a real misconfiguration should fail fast.
|
|
||||||
retry-max-attempts: 5
|
retry-max-attempts: 5
|
||||||
|
|
||||||
- name: Sign .exe files
|
- name: Sign .exe files
|
||||||
|
|||||||
Executable
+142
@@ -0,0 +1,142 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
#
|
||||||
|
# Backport a merged PR to a release branch.
|
||||||
|
#
|
||||||
|
# Cherry-picks the merge commit of a merged PR onto release-1.11 and opens a
|
||||||
|
# backport PR, mirroring the shape of #1842.
|
||||||
|
#
|
||||||
|
# Usage: ./backport.sh <pr-number> [target-version] [--continue]
|
||||||
|
# pr-number The merged PR to backport
|
||||||
|
# target-version Release version to target (default: 1.11)
|
||||||
|
# --continue Resume after resolving cherry-pick conflicts: assumes the
|
||||||
|
# fixes are staged, runs `git cherry-pick --continue`, and
|
||||||
|
# proceeds to push and open the backport PR.
|
||||||
|
#
|
||||||
|
# Requires: gh (authenticated), git, jq.
|
||||||
|
|
||||||
|
set -euo pipefail
|
||||||
|
|
||||||
|
PR_NUMBER=${1:?usage: backport.sh <pr-number> [target-version] [--continue]}
|
||||||
|
shift
|
||||||
|
CONTINUE=0
|
||||||
|
TARGET_VERSION=""
|
||||||
|
for arg in "$@"; do
|
||||||
|
case "$arg" in
|
||||||
|
--continue) CONTINUE=1 ;;
|
||||||
|
*) TARGET_VERSION=$arg ;;
|
||||||
|
esac
|
||||||
|
done
|
||||||
|
|
||||||
|
for cmd in gh git jq; do
|
||||||
|
command -v "$cmd" >/dev/null || { echo "error: $cmd is required" >&2; exit 1; }
|
||||||
|
done
|
||||||
|
|
||||||
|
# Pull the source PR's metadata. Refuse to backport a PR that never merged.
|
||||||
|
pr_json=$(gh pr view "$PR_NUMBER" --json title,body,mergedAt,mergeCommit,milestone)
|
||||||
|
PR_TITLE=$(jq -r '.title' <<<"$pr_json")
|
||||||
|
PR_BODY=$(jq -r '.body' <<<"$pr_json")
|
||||||
|
MERGE_SHA=$(jq -r '.mergeCommit.oid // empty' <<<"$pr_json")
|
||||||
|
PR_MILESTONE=$(jq -r '.milestone.title // empty' <<<"$pr_json")
|
||||||
|
|
||||||
|
if [ "$(jq -r '.mergedAt // empty' <<<"$pr_json")" = "" ] || [ -z "$MERGE_SHA" ]; then
|
||||||
|
echo "error: PR #${PR_NUMBER} is not merged (no merge commit to cherry-pick)" >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
|
# Default the target to one minor below the PR's milestone: a PR landing in
|
||||||
|
# v1.12.0 backports to release-1.11. An explicit target-version arg overrides.
|
||||||
|
if [ -z "$TARGET_VERSION" ]; then
|
||||||
|
if [[ $PR_MILESTONE =~ ^v?([0-9]+)\.([0-9]+) ]]; then
|
||||||
|
TARGET_VERSION="${BASH_REMATCH[1]}.$(( BASH_REMATCH[2] - 1 ))"
|
||||||
|
else
|
||||||
|
echo "error: PR #${PR_NUMBER} has no v<major>.<minor>.* milestone to infer the target from." >&2
|
||||||
|
echo "Pass the target version explicitly, e.g. $0 ${PR_NUMBER} 1.11" >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
fi
|
||||||
|
TARGET_BRANCH="release-${TARGET_VERSION}"
|
||||||
|
|
||||||
|
# Slug from the PR title: lower-case, non-alphanumerics to dashes, trimmed.
|
||||||
|
slug=$(printf '%s' "$PR_TITLE" \
|
||||||
|
| tr '[:upper:]' '[:lower:]' \
|
||||||
|
| sed -E 's/[^a-z0-9]+/-/g; s/^-+//; s/-+$//' \
|
||||||
|
| cut -c1-50 \
|
||||||
|
| sed -E 's/-+$//')
|
||||||
|
BRANCH="backport-${TARGET_VERSION//./-}-${slug}"
|
||||||
|
|
||||||
|
if [ "$CONTINUE" -eq 1 ]; then
|
||||||
|
# Resuming: the branch already exists and a cherry-pick is mid-conflict.
|
||||||
|
current=$(git rev-parse --abbrev-ref HEAD)
|
||||||
|
if [ "$current" != "$BRANCH" ]; then
|
||||||
|
echo "error: --continue expects to be on ${BRANCH}, but HEAD is ${current}" >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
echo "Resuming backport of #${PR_NUMBER} on ${BRANCH}"
|
||||||
|
# Conflicts assumed resolved and staged; core.editor=true keeps the commit message.
|
||||||
|
git -c core.editor=true cherry-pick --continue
|
||||||
|
else
|
||||||
|
# A fresh run switches branches and cherry-picks, so tracked changes must be
|
||||||
|
# clean. Untracked files are fine, and --continue is exempt (its resolved
|
||||||
|
# conflicts are meant to be staged).
|
||||||
|
if [ -n "$(git status --porcelain --untracked-files=no)" ]; then
|
||||||
|
echo "error: working tree has uncommitted changes to tracked files; commit or stash them first." >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
|
echo "Backporting #${PR_NUMBER} (${MERGE_SHA}) onto ${TARGET_BRANCH} as ${BRANCH}"
|
||||||
|
|
||||||
|
git fetch origin
|
||||||
|
git checkout -b "$BRANCH" "origin/${TARGET_BRANCH}"
|
||||||
|
|
||||||
|
# -m 1 handles a real merge commit; plain cherry-pick handles a squash merge.
|
||||||
|
if ! { git cherry-pick -x -m 1 "$MERGE_SHA" 2>/dev/null || git cherry-pick -x "$MERGE_SHA"; }; then
|
||||||
|
echo "error: ${MERGE_SHA} did not cherry-pick cleanly onto ${TARGET_BRANCH}." >&2
|
||||||
|
echo "Resolve the conflicts, 'git add' them, then re-run:" >&2
|
||||||
|
echo " $0 ${PR_NUMBER} ${TARGET_VERSION} --continue" >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
fi
|
||||||
|
|
||||||
|
title="backport v${TARGET_VERSION}: ${PR_TITLE}"
|
||||||
|
body=$(printf 'Backport from #%s to release v%s\n\n---\n\n%s' \
|
||||||
|
"$PR_NUMBER" "$TARGET_VERSION" "$PR_BODY")
|
||||||
|
|
||||||
|
# Highest open milestone matching v<version>.* (e.g. v1.11.1), if any.
|
||||||
|
MILESTONE=$(gh api "repos/{owner}/{repo}/milestones?state=open" \
|
||||||
|
--jq ".[].title | select(startswith(\"v${TARGET_VERSION}.\"))" \
|
||||||
|
| sort -V | tail -n1)
|
||||||
|
milestone_args=()
|
||||||
|
if [ -n "$MILESTONE" ]; then
|
||||||
|
milestone_args=(--milestone "$MILESTONE")
|
||||||
|
else
|
||||||
|
echo "warning: no open milestone matching v${TARGET_VERSION}.* found; PR will have no milestone" >&2
|
||||||
|
fi
|
||||||
|
|
||||||
|
echo
|
||||||
|
echo "About to run:"
|
||||||
|
printf ' git push -u --force-with-lease origin %q\n' "$BRANCH"
|
||||||
|
printf ' gh pr create --base %q --head %q --title %q --body %q' \
|
||||||
|
"$TARGET_BRANCH" "$BRANCH" "$title" "$body"
|
||||||
|
[ -n "$MILESTONE" ] && printf ' --milestone %q' "$MILESTONE"
|
||||||
|
printf '\n'
|
||||||
|
echo
|
||||||
|
read -r -p "Open this pull request? [y/N] " reply
|
||||||
|
case "$reply" in
|
||||||
|
[yY] | [yY][eE][sS]) ;;
|
||||||
|
*) echo "Aborted. The cherry-pick is on local branch ${BRANCH}; push it manually if desired."; exit 0 ;;
|
||||||
|
esac
|
||||||
|
|
||||||
|
git push -u --force-with-lease origin "$BRANCH"
|
||||||
|
|
||||||
|
gh pr create \
|
||||||
|
--base "$TARGET_BRANCH" \
|
||||||
|
--head "$BRANCH" \
|
||||||
|
--title "$title" \
|
||||||
|
--body "$body" \
|
||||||
|
${milestone_args[@]+"${milestone_args[@]}"}
|
||||||
|
|
||||||
|
# The backport now exists, so drop the label that flagged this PR for one.
|
||||||
|
gh pr edit "$PR_NUMBER" --remove-label needs-backport
|
||||||
|
|
||||||
|
# Switch back to the original branch we were on before doing the backport
|
||||||
|
git checkout -
|
||||||
+31
-32
@@ -20,45 +20,44 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v7
|
||||||
|
|
||||||
- uses: actions/setup-go@v7
|
|
||||||
with:
|
|
||||||
go-version: '1.26'
|
|
||||||
check-latest: true
|
|
||||||
|
|
||||||
- name: Smoke Docker
|
|
||||||
run: make smoke-docker
|
|
||||||
|
|
||||||
- name: Smoke Docker IPv6 overlay
|
|
||||||
run: make smoke-docker-ipv6
|
|
||||||
|
|
||||||
- name: Smoke Relay Docker
|
|
||||||
run: make smoke-relay-docker
|
|
||||||
|
|
||||||
- name: Smoke Docker boringcrypto
|
|
||||||
run: make boringcrypto smoke-docker
|
|
||||||
|
|
||||||
- name: Smoke Docker fips140
|
|
||||||
run: make fips140-all GOALS=smoke-docker
|
|
||||||
|
|
||||||
timeout-minutes: 10
|
|
||||||
|
|
||||||
smoke-self:
|
|
||||||
name: Run self traffic smoke test on macOS
|
|
||||||
runs-on: macos-latest
|
|
||||||
steps:
|
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
|
||||||
|
|
||||||
- uses: actions/setup-go@v7
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: '1.26'
|
go-version: '1.26'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: build
|
- name: build
|
||||||
run: make bin
|
run: make bin-docker CGO_ENABLED=1 BUILD_ARGS=-race
|
||||||
|
|
||||||
- name: run smoke-self
|
- name: setup docker image
|
||||||
working-directory: ./.github/workflows/smoke
|
working-directory: ./.github/workflows/smoke
|
||||||
run: ./smoke-self.sh
|
run: ./build.sh
|
||||||
|
|
||||||
|
- name: run smoke
|
||||||
|
working-directory: ./.github/workflows/smoke
|
||||||
|
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
|
||||||
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: ./build-relay.sh
|
||||||
|
|
||||||
|
- name: run smoke relay
|
||||||
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: ./smoke-relay.sh
|
||||||
|
|
||||||
|
- name: setup docker image for P256
|
||||||
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: NAME="smoke-p256" CURVE=P256 ./build.sh
|
||||||
|
|
||||||
|
- name: run smoke-p256
|
||||||
|
working-directory: ./.github/workflows/smoke
|
||||||
|
run: NAME="smoke-p256" ./smoke.sh
|
||||||
|
|
||||||
timeout-minutes: 10
|
timeout-minutes: 10
|
||||||
|
|||||||
@@ -1,130 +0,0 @@
|
|||||||
#!/bin/bash
|
|
||||||
|
|
||||||
# A host must be able to reach its own overlay address. Where the kernel sends
|
|
||||||
# that traffic through the tun rather than over loopback, nebula sees it and
|
|
||||||
# hands it straight back (immediatelyForwardToSelf), and whether the kernel
|
|
||||||
# accepts what comes back is only answerable against a real kernel. Runs one
|
|
||||||
# nebula on this machine as root and aims every probe at its own address.
|
|
||||||
|
|
||||||
set -e -x
|
|
||||||
|
|
||||||
set -o pipefail
|
|
||||||
|
|
||||||
V4=192.0.2.1
|
|
||||||
V6=2001:db8::1
|
|
||||||
|
|
||||||
case "$(uname -s)" in
|
|
||||||
Darwin) TUN_DEV=utun ;;
|
|
||||||
*) TUN_DEV=tun0 ;;
|
|
||||||
esac
|
|
||||||
|
|
||||||
ROOT="$(cd ../../.. && pwd)"
|
|
||||||
|
|
||||||
rm -rf build/self
|
|
||||||
mkdir -p build/self
|
|
||||||
cd build/self
|
|
||||||
|
|
||||||
cleanup() {
|
|
||||||
echo
|
|
||||||
echo " *** cleanup"
|
|
||||||
echo
|
|
||||||
|
|
||||||
set +e
|
|
||||||
if [ -n "$NEBULA_PID" ]
|
|
||||||
then
|
|
||||||
sudo kill "$NEBULA_PID"
|
|
||||||
fi
|
|
||||||
{ kill $(jobs -p); wait; } 2>/dev/null
|
|
||||||
sed 's/^/ [self] /' nebula.log
|
|
||||||
}
|
|
||||||
|
|
||||||
trap cleanup EXIT
|
|
||||||
|
|
||||||
# perl is on every platform this runs on; timeout(1) is not.
|
|
||||||
alarm() {
|
|
||||||
perl -e 'alarm shift; exec @ARGV' "$@"
|
|
||||||
}
|
|
||||||
|
|
||||||
RESULTS=""
|
|
||||||
FAILED=""
|
|
||||||
probe() {
|
|
||||||
local name="$1"
|
|
||||||
shift
|
|
||||||
if "$@"
|
|
||||||
then
|
|
||||||
RESULTS="$RESULTS $name=ok"
|
|
||||||
else
|
|
||||||
RESULTS="$RESULTS $name=FAIL"
|
|
||||||
FAILED="$FAILED $name"
|
|
||||||
fi
|
|
||||||
}
|
|
||||||
|
|
||||||
# Send one datagram, then wait for the listener to have written it out.
|
|
||||||
udp_probe() {
|
|
||||||
echo self | alarm 5 nc -u -w1 "$1" 3000 || true
|
|
||||||
set +x
|
|
||||||
for _ in $(seq 1 20)
|
|
||||||
do
|
|
||||||
if grep -q self "$2"
|
|
||||||
then
|
|
||||||
set -x
|
|
||||||
return 0
|
|
||||||
fi
|
|
||||||
sleep 0.25
|
|
||||||
done
|
|
||||||
set -x
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
|
|
||||||
"$ROOT/nebula-cert" ca -name "Smoke Test"
|
|
||||||
"$ROOT/nebula-cert" sign -name self -networks "$V4/24,$V6/64"
|
|
||||||
|
|
||||||
HOST=self AM_LIGHTHOUSE=true TUN_DEV="$TUN_DEV" ../../genconfig.sh >self.yml
|
|
||||||
|
|
||||||
"$ROOT/nebula" -config self.yml -test
|
|
||||||
|
|
||||||
sudo -v
|
|
||||||
sudo "$ROOT/nebula" -config self.yml >nebula.log 2>&1 &
|
|
||||||
NEBULA_PID=$!
|
|
||||||
|
|
||||||
for _ in $(seq 1 40)
|
|
||||||
do
|
|
||||||
ifconfig | grep "inet6 $V6 " >/dev/null && break
|
|
||||||
sleep 0.25
|
|
||||||
done
|
|
||||||
ifconfig | grep "inet $V4 "
|
|
||||||
ifconfig | grep "inet6 $V6 "
|
|
||||||
|
|
||||||
nc -l "$V4" 2000 >/dev/null &
|
|
||||||
nc -l "$V6" 2000 >/dev/null &
|
|
||||||
nc -u -l "$V4" 3000 >udp4.txt &
|
|
||||||
nc -u -l "$V6" 3000 >udp6.txt &
|
|
||||||
sleep 1
|
|
||||||
|
|
||||||
set +x
|
|
||||||
echo
|
|
||||||
echo " *** Testing self traffic from $V4"
|
|
||||||
echo
|
|
||||||
set -x
|
|
||||||
probe icmp4 alarm 5 ping -c1 "$V4"
|
|
||||||
probe tcp4 alarm 5 nc -z "$V4" 2000
|
|
||||||
probe udp4 udp_probe "$V4" udp4.txt
|
|
||||||
|
|
||||||
set +x
|
|
||||||
echo
|
|
||||||
echo " *** Testing self traffic from $V6"
|
|
||||||
echo
|
|
||||||
set -x
|
|
||||||
probe icmp6 alarm 5 ping6 -c1 "$V6"
|
|
||||||
probe tcp6 alarm 5 nc -z "$V6" 2000
|
|
||||||
probe udp6 udp_probe "$V6" udp6.txt
|
|
||||||
|
|
||||||
set +x
|
|
||||||
echo
|
|
||||||
echo " *** self traffic:$RESULTS"
|
|
||||||
echo
|
|
||||||
if [ -n "$FAILED" ]
|
|
||||||
then
|
|
||||||
echo "self traffic failed:$FAILED" >&2
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
@@ -51,19 +51,15 @@ wsl -d $Distro -- bash -c "rm -rf $WslDir && mkdir -p $WslDir" | Out-Null
|
|||||||
$DevName = 'nebula-smoke'
|
$DevName = 'nebula-smoke'
|
||||||
$Ip1 = '192.168.241.1'
|
$Ip1 = '192.168.241.1'
|
||||||
$Ip2 = '192.168.241.2'
|
$Ip2 = '192.168.241.2'
|
||||||
# Dual stack on purpose: a v4-only overlay never exercises the v6 side of tun.mtu.
|
|
||||||
$Ip6_1 = 'fd42:4242:241::1'
|
|
||||||
$Ip6_2 = 'fd42:4242:241::2'
|
|
||||||
$Mtu = 1300
|
|
||||||
$Port = 4242
|
$Port = 4242
|
||||||
|
|
||||||
& $NebulaCert ca -name 'smoke-ca' -out-crt "$WorkDir\ca.crt" -out-key "$WorkDir\ca.key"
|
& $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)" }
|
if ($LASTEXITCODE -ne 0) { throw "nebula-cert ca failed (exit $LASTEXITCODE)" }
|
||||||
|
|
||||||
& $NebulaCert sign -name 'lighthouse' -networks "$Ip1/24,$Ip6_1/64" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\lighthouse.crt" -out-key "$WorkDir\lighthouse.key"
|
& $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)" }
|
if ($LASTEXITCODE -ne 0) { throw "nebula-cert sign lighthouse failed (exit $LASTEXITCODE)" }
|
||||||
|
|
||||||
& $NebulaCert sign -name 'peer' -networks "$Ip2/24,$Ip6_2/64" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\peer.crt" -out-key "$WorkDir\peer.key"
|
& $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)" }
|
if ($LASTEXITCODE -ne 0) { throw "nebula-cert sign peer failed (exit $LASTEXITCODE)" }
|
||||||
|
|
||||||
# Windows lighthouse config.
|
# Windows lighthouse config.
|
||||||
@@ -86,7 +82,7 @@ tun:
|
|||||||
drop_local_broadcast: false
|
drop_local_broadcast: false
|
||||||
drop_multicast: false
|
drop_multicast: false
|
||||||
tx_queue: 500
|
tx_queue: 500
|
||||||
mtu: $Mtu
|
mtu: 1300
|
||||||
network_category: private
|
network_category: private
|
||||||
logging:
|
logging:
|
||||||
level: info
|
level: info
|
||||||
@@ -130,7 +126,7 @@ tun:
|
|||||||
drop_local_broadcast: false
|
drop_local_broadcast: false
|
||||||
drop_multicast: false
|
drop_multicast: false
|
||||||
tx_queue: 500
|
tx_queue: 500
|
||||||
mtu: $Mtu
|
mtu: 1300
|
||||||
logging:
|
logging:
|
||||||
level: info
|
level: info
|
||||||
format: text
|
format: text
|
||||||
@@ -173,7 +169,7 @@ Write-Host '=== WSL diagnostic ==='
|
|||||||
wsl --version 2>&1 | Out-Host
|
wsl --version 2>&1 | Out-Host
|
||||||
wsl --list --verbose 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 -- 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; { echo 0 > /proc/sys/net/ipv6/conf/all/disable_ipv6; echo 0 > /proc/sys/net/ipv6/conf/default/disable_ipv6; } 2>/dev/null || true; ls -l /dev/net/tun"
|
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?)" }
|
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
|
# Deliberately no New-NetFirewallRule calls here -- nebula's windows_bypass_wdf
|
||||||
@@ -218,16 +214,6 @@ try {
|
|||||||
}
|
}
|
||||||
Write-Host "OK: $DevName NetworkCategory=Private"
|
Write-Host "OK: $DevName NetworkCategory=Private"
|
||||||
|
|
||||||
# v6 silently kept the adapter default of 65535 while v4 was correct.
|
|
||||||
foreach ($family in @('IPv4', 'IPv6')) {
|
|
||||||
Wait-Until -TimeoutSec 30 -What "$DevName $family NlMtu=$Mtu" -Predicate {
|
|
||||||
if ($lhProc.HasExited) { throw "lighthouse exited (code $($lhProc.ExitCode)) before $family mtu was set" }
|
|
||||||
$rows = @(Get-NetIPInterface -InterfaceAlias $DevName -AddressFamily $family -ErrorAction SilentlyContinue)
|
|
||||||
$rows.Count -gt 0 -and -not ($rows | Where-Object { $_.NlMtu -ne $Mtu })
|
|
||||||
}
|
|
||||||
Write-Host "OK: $DevName $family NlMtu=$Mtu"
|
|
||||||
}
|
|
||||||
|
|
||||||
Wait-Until -TimeoutSec 30 -What "WSL nebula1 with $Ip2" -Predicate {
|
Wait-Until -TimeoutSec 30 -What "WSL nebula1 with $Ip2" -Predicate {
|
||||||
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before tun was ready" }
|
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 = wsl -d $Distro -u root -- bash -c "ip -o addr show nebula1 2>/dev/null | grep -q 'inet $Ip2' && echo yes"
|
||||||
@@ -235,13 +221,6 @@ try {
|
|||||||
}
|
}
|
||||||
Write-Host "OK: WSL nebula1 has $Ip2"
|
Write-Host "OK: WSL nebula1 has $Ip2"
|
||||||
|
|
||||||
Wait-Until -TimeoutSec 30 -What "WSL nebula1 with $Ip6_2" -Predicate {
|
|
||||||
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before the v6 address was up" }
|
|
||||||
$r = wsl -d $Distro -u root -- bash -c "ip -o addr show nebula1 2>/dev/null | grep -q 'inet6 $Ip6_2' && echo yes"
|
|
||||||
("$r").Trim() -eq 'yes'
|
|
||||||
}
|
|
||||||
Write-Host "OK: WSL nebula1 has $Ip6_2"
|
|
||||||
|
|
||||||
Wait-Until -TimeoutSec 30 -What "ping from WSL peer to windows lighthouse ($Ip1)" -Predicate {
|
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" }
|
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 = wsl -d $Distro -u root -- bash -c "ping -c1 -W1 $Ip1 >/dev/null 2>&1 && echo OK"
|
||||||
@@ -255,14 +234,6 @@ try {
|
|||||||
}
|
}
|
||||||
Write-Host "OK: windows lighthouse -> WSL peer"
|
Write-Host "OK: windows lighthouse -> WSL peer"
|
||||||
|
|
||||||
# Otherwise the v6 networks only prove the interface exists, not that it forwards.
|
|
||||||
Wait-Until -TimeoutSec 30 -What "v6 ping from WSL peer to windows lighthouse ($Ip6_1)" -Predicate {
|
|
||||||
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before the v6 ping succeeded" }
|
|
||||||
$r = wsl -d $Distro -u root -- bash -c "ping -6 -c1 -W1 $Ip6_1 >/dev/null 2>&1 && echo OK"
|
|
||||||
("$r").Trim() -eq 'OK'
|
|
||||||
}
|
|
||||||
Write-Host "OK: WSL peer -> windows lighthouse over v6"
|
|
||||||
|
|
||||||
Write-Host ''
|
Write-Host ''
|
||||||
Write-Host 'All smoke checks passed.'
|
Write-Host 'All smoke checks passed.'
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -58,14 +58,9 @@ jobs:
|
|||||||
e2e-cmd: make e2evv
|
e2e-cmd: make e2evv
|
||||||
- name: linux-boringcrypto
|
- name: linux-boringcrypto
|
||||||
os: ubuntu-latest
|
os: ubuntu-latest
|
||||||
build-cmd: make boringcrypto
|
build-cmd: make bin-boringcrypto
|
||||||
test-cmd: make boringcrypto test
|
test-cmd: make test-boringcrypto
|
||||||
e2e-cmd: make boringcrypto e2evv
|
e2e-cmd: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0"
|
||||||
- name: linux-fips140
|
|
||||||
os: ubuntu-latest
|
|
||||||
build-cmd: make fips140-all
|
|
||||||
test-cmd: make fips140-all GOALS=test
|
|
||||||
e2e-cmd: make fips140-all GOALS=e2evv
|
|
||||||
- name: linux-pkcs11
|
- name: linux-pkcs11
|
||||||
os: ubuntu-latest
|
os: ubuntu-latest
|
||||||
build-cmd: make bin-pkcs11
|
build-cmd: make bin-pkcs11
|
||||||
|
|||||||
+1
-28
@@ -7,31 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
## [Unreleased]
|
## [Unreleased]
|
||||||
|
|
||||||
## [1.11.1] - 2026-08-21
|
|
||||||
|
|
||||||
See the [v1.11.1](https://github.com/slackhq/nebula/milestone/30?closed=1) milestone for a complete list of changes.
|
|
||||||
|
|
||||||
### Changed
|
|
||||||
|
|
||||||
- IPv6 packets whose next header is a protocol Nebula does not parse (SCTP, GRE, IP-in-IP, etc.) are now
|
|
||||||
classified as that protocol with no ports, closing a firewall bypass where a crafted payload could steer
|
|
||||||
the classifier into reading one as TCP/UDP and matching a TCP/UDP rule. These packets are now matched as
|
|
||||||
their true protocol, so only a `proto: any` rule allows them. If you carry one of these protocols over the
|
|
||||||
overlay, confirm a `proto: any` rule covers it before upgrading, it may have been passing only through this
|
|
||||||
bypass. (#1840)
|
|
||||||
- Drop the dependency on `github.com/cyberdelia/go-metrics-graphite`, which has been unmaintained for over ten
|
|
||||||
years, by inlining the small amount of code Nebula used. (#1832)
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- The ICMPv6 type was read from the wrong byte when classifying IPv6 packets, so the echo identifier used
|
|
||||||
for conntrack was never picked up. (#1840)
|
|
||||||
- Enforce outbound message counter limits so a tunnel is rehandshaked before the counter can wrap, preventing
|
|
||||||
nonce reuse. This is unreachable in practice, but is enforced as a defense-in-depth measure. (#1841)
|
|
||||||
- Prevent `nebula-cert ca` from running out of memory on 32bit systems when generating encrypted private keys. (#1834)
|
|
||||||
- Tolerate `ErrDumpInterrupted` when listing tun addresses on Linux, so a transient interrupted netlink dump
|
|
||||||
no longer aborts startup. (#1835)
|
|
||||||
|
|
||||||
## [1.11.0] - 2026-07-23
|
## [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.
|
See the [v1.11.0](https://github.com/slackhq/nebula/milestone/25?closed=1) milestone for a complete list of changes.
|
||||||
@@ -895,9 +870,7 @@ created.)
|
|||||||
|
|
||||||
- Initial public release.
|
- Initial public release.
|
||||||
|
|
||||||
[Unreleased]: https://github.com/slackhq/nebula/compare/v1.11.1...HEAD
|
[Unreleased]: https://github.com/slackhq/nebula/compare/v1.10.3...HEAD
|
||||||
[1.11.1]: https://github.com/slackhq/nebula/releases/tag/v1.11.1
|
|
||||||
[1.11.0]: https://github.com/slackhq/nebula/releases/tag/v1.11.0
|
|
||||||
[1.10.3]: https://github.com/slackhq/nebula/releases/tag/v1.10.3
|
[1.10.3]: https://github.com/slackhq/nebula/releases/tag/v1.10.3
|
||||||
[1.10.2]: https://github.com/slackhq/nebula/releases/tag/v1.10.2
|
[1.10.2]: https://github.com/slackhq/nebula/releases/tag/v1.10.2
|
||||||
[1.10.1]: https://github.com/slackhq/nebula/releases/tag/v1.10.1
|
[1.10.1]: https://github.com/slackhq/nebula/releases/tag/v1.10.1
|
||||||
|
|||||||
@@ -72,17 +72,6 @@ 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_MIPS = linux-mips linux-mipsle linux-mips64 linux-mips64le linux-mips-softfloat
|
||||||
ALL_CROSS_LINUX_OTHER = linux-386 linux-ppc64le linux-riscv64 linux-loong64
|
ALL_CROSS_LINUX_OTHER = linux-386 linux-ppc64le linux-riscv64 linux-loong64
|
||||||
|
|
||||||
# Based on section 2.2 of the Go Cryptographic Module CVMP Security Policy #5247
|
|
||||||
ALL_FIPS140 = linux-amd64-fips140 \
|
|
||||||
linux-arm64-fips140 \
|
|
||||||
windows-amd64-fips140 \
|
|
||||||
windows-arm64-fips140 \
|
|
||||||
darwin-arm64-fips140 \
|
|
||||||
freebsd-amd64-fips140 \
|
|
||||||
linux-arm-7-fips140 \
|
|
||||||
linux-mips64-fips140 \
|
|
||||||
linux-ppc64le-fips140
|
|
||||||
|
|
||||||
e2e:
|
e2e:
|
||||||
$(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e
|
$(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e
|
||||||
|
|
||||||
@@ -148,8 +137,6 @@ release-netbsd: $(ALL_NETBSD:%=build/nebula-%.tar.gz)
|
|||||||
|
|
||||||
release-boringcrypto: build/nebula-linux-$(shell go env GOARCH)-boringcrypto.tar.gz
|
release-boringcrypto: build/nebula-linux-$(shell go env GOARCH)-boringcrypto.tar.gz
|
||||||
|
|
||||||
release-fips140: $(ALL_FIPS140:%=build/nebula-%.tar.gz)
|
|
||||||
|
|
||||||
BUILD_ARGS += -trimpath
|
BUILD_ARGS += -trimpath
|
||||||
|
|
||||||
bin-windows: build/windows-amd64/nebula.exe build/windows-amd64/nebula-cert.exe
|
bin-windows: build/windows-amd64/nebula.exe build/windows-amd64/nebula-cert.exe
|
||||||
@@ -170,24 +157,17 @@ bin-freebsd-arm64: build/freebsd-arm64/nebula build/freebsd-arm64/nebula-cert
|
|||||||
bin-boringcrypto: build/linux-$(shell go env GOARCH)-boringcrypto/nebula build/linux-$(shell go env GOARCH)-boringcrypto/nebula-cert
|
bin-boringcrypto: build/linux-$(shell go env GOARCH)-boringcrypto/nebula build/linux-$(shell go env GOARCH)-boringcrypto/nebula-cert
|
||||||
mv $? .
|
mv $? .
|
||||||
|
|
||||||
bin-fips140: build/linux-$(shell go env GOARCH)-fips140/nebula build/linux-$(shell go env GOARCH)-fips140/nebula-cert
|
|
||||||
mv $? .
|
|
||||||
|
|
||||||
bin-pkcs11: BUILD_ARGS += -tags pkcs11
|
bin-pkcs11: BUILD_ARGS += -tags pkcs11
|
||||||
bin-pkcs11: CGO_ENABLED = 1
|
bin-pkcs11: CGO_ENABLED = 1
|
||||||
bin-pkcs11: bin
|
bin-pkcs11: bin
|
||||||
|
|
||||||
# Build with the pprof debug server (serves on :6060). See startPprofServer.
|
|
||||||
debug: BUILD_ARGS += -tags debug
|
|
||||||
debug: bin
|
|
||||||
|
|
||||||
bin:
|
bin:
|
||||||
$(GOENV) go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
||||||
$(GOENV) go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
||||||
|
|
||||||
install:
|
install:
|
||||||
$(GOENV) go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ${NEBULA_CMD_PATH}
|
go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ${NEBULA_CMD_PATH}
|
||||||
$(GOENV) go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ./cmd/nebula-cert
|
go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ./cmd/nebula-cert
|
||||||
|
|
||||||
build/linux-arm-%: GOENV += GOARM=$(word 3, $(subst -, ,$*))
|
build/linux-arm-%: GOENV += GOARM=$(word 3, $(subst -, ,$*))
|
||||||
build/linux-mips-%: GOENV += GOMIPS=$(word 3, $(subst -, ,$*))
|
build/linux-mips-%: GOENV += GOMIPS=$(word 3, $(subst -, ,$*))
|
||||||
@@ -198,11 +178,8 @@ build/linux-mips-softfloat/%: LDFLAGS += -s -w
|
|||||||
# boringcrypto
|
# boringcrypto
|
||||||
build/linux-amd64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
build/linux-amd64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
||||||
build/linux-arm64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
build/linux-arm64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
||||||
|
build/linux-amd64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
||||||
# fips140
|
build/linux-arm64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
||||||
FIPSVERSION = v1.0.0
|
|
||||||
$(foreach _rule, $(ALL_FIPS140), build/$(_rule)/%): GOENV += GOFIPS140=$(FIPSVERSION)
|
|
||||||
$(foreach _rule, $(ALL_FIPS140), build/$(_rule)/%): BUILD_ARGS += -tags fips140enforce
|
|
||||||
|
|
||||||
build/%/nebula: .FORCE
|
build/%/nebula: .FORCE
|
||||||
GOOS=$(firstword $(subst -, , $*)) \
|
GOOS=$(firstword $(subst -, , $*)) \
|
||||||
@@ -233,7 +210,10 @@ vet:
|
|||||||
go vet $(VET_FLAGS) -v ./...
|
go vet $(VET_FLAGS) -v ./...
|
||||||
|
|
||||||
test:
|
test:
|
||||||
$(TEST_ENV) go test $(TEST_FLAGS) -v ./...
|
go test -v ./...
|
||||||
|
|
||||||
|
test-boringcrypto:
|
||||||
|
GOEXPERIMENT=boringcrypto CGO_ENABLED=1 go test -ldflags "-checklinkname=0" -v ./...
|
||||||
|
|
||||||
test-pkcs11:
|
test-pkcs11:
|
||||||
CGO_ENABLED=1 go test -v -tags pkcs11 ./...
|
CGO_ENABLED=1 go test -v -tags pkcs11 ./...
|
||||||
@@ -276,75 +256,29 @@ ifeq ($(words $(MAKECMDGOALS)),1)
|
|||||||
@$(MAKE) service ${.DEFAULT_GOAL} --no-print-directory
|
@$(MAKE) service ${.DEFAULT_GOAL} --no-print-directory
|
||||||
endif
|
endif
|
||||||
|
|
||||||
# Useful to chain together, like:
|
|
||||||
# - make fips140 e2evv
|
|
||||||
# - make fips140 smoke-docker
|
|
||||||
# Use `release-fips140` to build release binaries
|
|
||||||
fips140:
|
|
||||||
@echo > $(NULL_FILE)
|
|
||||||
ifeq ($(strip $(GOFIPS140)),)
|
|
||||||
$(eval GOFIPS140 = $(FIPSVERSION))
|
|
||||||
endif
|
|
||||||
$(eval GOENV += GOFIPS140=$(GOFIPS140))
|
|
||||||
$(eval BUILD_ARGS += -tags fips140enforce)
|
|
||||||
$(eval TEST_ENV += $(GOENV))
|
|
||||||
$(eval CURVE = P256)
|
|
||||||
ifeq ($(words $(MAKECMDGOALS)),1)
|
|
||||||
@$(MAKE) fips140 GOFIPS140=$(GOFIPS140) ${.DEFAULT_GOAL} --no-print-directory
|
|
||||||
endif
|
|
||||||
|
|
||||||
# To test the future pending module, use like `make fips140-latest test`
|
|
||||||
ALL_GOFIPS140 = v1.0.0 v1.26.0 latest
|
|
||||||
define FIPS140_rule
|
|
||||||
fips140-$(1): GOFIPS140 = $(1)
|
|
||||||
fips140-$(1): fips140
|
|
||||||
endef
|
|
||||||
$(foreach _rule, $(ALL_GOFIPS140), $(eval $(call FIPS140_rule,$(_rule))))
|
|
||||||
|
|
||||||
# Iterate and run the goals for all fips versions, like `make fips140-all GOALS=test`
|
|
||||||
fips140-all:
|
|
||||||
@$(foreach _v,$(ALL_GOFIPS140),$(MAKE) fips140-$(_v) $(GOALS) &&) true
|
|
||||||
|
|
||||||
# Useful to chain together, like:
|
|
||||||
# - make boringcrypto e2evv
|
|
||||||
# - make boringcrypto smoke-docker
|
|
||||||
# Use `release-boringcrypto` or `bin-boringcrypto` to build release binaries
|
|
||||||
boringcrypto:
|
|
||||||
@echo > $(NULL_FILE)
|
|
||||||
$(eval GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1)
|
|
||||||
$(eval TEST_ENV += $(GOENV))
|
|
||||||
$(eval CURVE = P256)
|
|
||||||
ifeq ($(words $(MAKECMDGOALS)),1)
|
|
||||||
@$(MAKE) boringcrypto ${.DEFAULT_GOAL} --no-print-directory
|
|
||||||
endif
|
|
||||||
|
|
||||||
bin-docker: bin build/linux-amd64/nebula build/linux-amd64/nebula-cert
|
bin-docker: bin build/linux-amd64/nebula build/linux-amd64/nebula-cert
|
||||||
|
|
||||||
smoke-docker: BUILD_ARGS += -race
|
|
||||||
smoke-docker: GOENV += CGO_ENABLED=1
|
|
||||||
smoke-docker: bin-docker
|
smoke-docker: bin-docker
|
||||||
# This is so we can limit `fips140` smoke test to just P256 curve.
|
cd .github/workflows/smoke/ && ./build.sh
|
||||||
if [ "$(CURVE)" != "P256" ]; then cd .github/workflows/smoke/ && $(GOENV) ./build.sh; fi
|
cd .github/workflows/smoke/ && ./smoke.sh
|
||||||
if [ "$(CURVE)" != "P256" ]; then cd .github/workflows/smoke/ && $(GOENV) ./smoke.sh; fi
|
cd .github/workflows/smoke/ && NAME="smoke-p256" CURVE="P256" ./build.sh
|
||||||
cd .github/workflows/smoke/ && $(GOENV) NAME="smoke-p256" CURVE="P256" ./build.sh
|
cd .github/workflows/smoke/ && NAME="smoke-p256" ./smoke.sh
|
||||||
cd .github/workflows/smoke/ && $(GOENV) NAME="smoke-p256" ./smoke.sh
|
|
||||||
|
|
||||||
smoke-relay-docker: BUILD_ARGS += -race
|
|
||||||
smoke-relay-docker: GOENV += CGO_ENABLED=1
|
|
||||||
smoke-relay-docker: bin-docker
|
smoke-relay-docker: bin-docker
|
||||||
cd .github/workflows/smoke/ && $(GOENV) ./build-relay.sh
|
cd .github/workflows/smoke/ && ./build-relay.sh
|
||||||
cd .github/workflows/smoke/ && $(GOENV) ./smoke-relay.sh
|
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
||||||
|
|
||||||
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
|
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
|
||||||
smoke-docker-ipv6: smoke-docker
|
smoke-docker-ipv6: smoke-docker
|
||||||
|
|
||||||
smoke-self: bin
|
smoke-docker-race: BUILD_ARGS = -race
|
||||||
cd .github/workflows/smoke/ && ./smoke-self.sh
|
smoke-docker-race: CGO_ENABLED = 1
|
||||||
|
smoke-docker-race: smoke-docker
|
||||||
|
|
||||||
smoke-vagrant/%: bin-docker build/%/nebula
|
smoke-vagrant/%: bin-docker build/%/nebula
|
||||||
cd .github/workflows/smoke/ && ./build.sh $*
|
cd .github/workflows/smoke/ && ./build.sh $*
|
||||||
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
||||||
|
|
||||||
.FORCE:
|
.FORCE:
|
||||||
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin bin-windows bin-windows-arm64 bin-darwin bin-freebsd bin-freebsd-arm64 bin-boringcrypto bin-fips140 bin-pkcs11 bin-docker boringcrypto build-test-mobile debug docker e2e e2ev e2evv e2evvv e2evvvv e2e-bench fips140 fips140-all $(ALL_GOFIPS140:%=fips140-%) install proto release release-linux release-freebsd release-openbsd release-netbsd release-boringcrypto release-fips140 service smoke-docker smoke-relay-docker smoke-docker-ipv6 smoke-self test test-pkcs11 test-cov-html vet smoke-vagrant/%
|
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
||||||
.DEFAULT_GOAL := bin
|
.DEFAULT_GOAL := bin
|
||||||
|
|||||||
@@ -145,27 +145,17 @@ To build nebula for a specific platform (ex, Windows):
|
|||||||
|
|
||||||
See the [Makefile](Makefile) for more details on build targets
|
See the [Makefile](Makefile) for more details on build targets
|
||||||
|
|
||||||
## Curve P256 and FIPS 140-3 mode
|
## Curve P256 and BoringCrypto
|
||||||
|
|
||||||
The default curve used for cryptographic handshakes and signatures is Curve25519. This is the recommended setting for most users. If your deployment has certain compliance requirements, you have the option of creating your CA using `nebula-cert ca -curve P256` to use NIST Curve P256. The CA will then sign certificates using ECDSA P256, and any hosts using these certificates will use P256 for ECDH handshakes.
|
The default curve used for cryptographic handshakes and signatures is Curve25519. This is the recommended setting for most users. If your deployment has certain compliance requirements, you have the option of creating your CA using `nebula-cert ca -curve P256` to use NIST Curve P256. The CA will then sign certificates using ECDSA P256, and any hosts using these certificates will use P256 for ECDH handshakes.
|
||||||
|
|
||||||
Nebula can be built to support the [FIPS 140-3](https://go.dev/doc/security/fips140) mode of Go by running either of the following make targets. (This sets GOFIPS140=v1.0.0, which must be done at compile time so that the correct AES-GCM can be used for FIPS 140-3 enforcement mode).
|
In addition, Nebula can be built using the [BoringCrypto GOEXPERIMENT](https://github.com/golang/go/blob/go1.20/src/crypto/internal/boring/README.md) by running either of the following make targets:
|
||||||
|
|
||||||
```sh
|
|
||||||
make fips140
|
|
||||||
make fips140 test
|
|
||||||
make release-fips140
|
|
||||||
```
|
|
||||||
|
|
||||||
Nebula can also be built using the [BoringCrypto GOEXPERIMENT](https://github.com/golang/go/blob/go1.20/src/crypto/internal/boring/README.md) by running either of the following make targets.
|
|
||||||
|
|
||||||
```sh
|
```sh
|
||||||
make bin-boringcrypto
|
make bin-boringcrypto
|
||||||
make release-boringcrypto
|
make release-boringcrypto
|
||||||
```
|
```
|
||||||
|
|
||||||
NOTE: boringcrypto support is deprecated and will be removed in the next release. Users should migrate to the native FIPS 140-3 mode described above.
|
|
||||||
|
|
||||||
This is not the recommended default deployment, but may be useful based on your compliance requirements.
|
This is not the recommended default deployment, but may be useful based on your compliance requirements.
|
||||||
|
|
||||||
## Credits
|
## Credits
|
||||||
|
|||||||
+3
-30
@@ -3,14 +3,11 @@ package main
|
|||||||
import (
|
import (
|
||||||
"crypto/ecdsa"
|
"crypto/ecdsa"
|
||||||
"crypto/elliptic"
|
"crypto/elliptic"
|
||||||
"crypto/fips140"
|
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"errors"
|
|
||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"math"
|
"math"
|
||||||
"math/bits"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -46,28 +43,7 @@ type caFlags struct {
|
|||||||
subnets *string
|
subnets *string
|
||||||
}
|
}
|
||||||
|
|
||||||
func defaultCurve() string {
|
|
||||||
if fips140.Enforced() {
|
|
||||||
return "P256"
|
|
||||||
}
|
|
||||||
return "25519"
|
|
||||||
}
|
|
||||||
|
|
||||||
func newCaFlags() *caFlags {
|
func newCaFlags() *caFlags {
|
||||||
// prevent running out of memory on 32-bit systems by defaulting to
|
|
||||||
// RFC9106's recommendation for memory-constrained environments
|
|
||||||
var (
|
|
||||||
defaultArgonMemory uint
|
|
||||||
defaultArgonIterations uint
|
|
||||||
)
|
|
||||||
if bits.UintSize == 32 {
|
|
||||||
defaultArgonMemory = 64 * 1024
|
|
||||||
defaultArgonIterations = 3
|
|
||||||
} else {
|
|
||||||
defaultArgonMemory = 2 * 1024 * 1024
|
|
||||||
defaultArgonIterations = 1
|
|
||||||
}
|
|
||||||
|
|
||||||
cf := caFlags{set: flag.NewFlagSet("ca", flag.ContinueOnError)}
|
cf := caFlags{set: flag.NewFlagSet("ca", flag.ContinueOnError)}
|
||||||
cf.set.Usage = func() {}
|
cf.set.Usage = func() {}
|
||||||
cf.name = cf.set.String("name", "", "Required: name of the certificate authority")
|
cf.name = cf.set.String("name", "", "Required: name of the certificate authority")
|
||||||
@@ -79,11 +55,11 @@ func newCaFlags() *caFlags {
|
|||||||
cf.groups = cf.set.String("groups", "", "Optional: comma separated list of groups. This will limit which groups subordinate certs can use")
|
cf.groups = cf.set.String("groups", "", "Optional: comma separated list of groups. This will limit which groups subordinate certs can use")
|
||||||
cf.networks = cf.set.String("networks", "", "Optional: comma separated list of ip address and network in CIDR notation. This will limit which ip addresses and networks subordinate certs can use in networks")
|
cf.networks = cf.set.String("networks", "", "Optional: comma separated list of ip address and network in CIDR notation. This will limit which ip addresses and networks subordinate certs can use in networks")
|
||||||
cf.unsafeNetworks = cf.set.String("unsafe-networks", "", "Optional: comma separated list of ip address and network in CIDR notation. This will limit which ip addresses and networks subordinate certs can use in unsafe networks")
|
cf.unsafeNetworks = cf.set.String("unsafe-networks", "", "Optional: comma separated list of ip address and network in CIDR notation. This will limit which ip addresses and networks subordinate certs can use in unsafe networks")
|
||||||
cf.argonMemory = cf.set.Uint("argon-memory", defaultArgonMemory, "Optional: Argon2 memory parameter (in KiB) used for encrypted private key passphrase")
|
cf.argonMemory = cf.set.Uint("argon-memory", 2*1024*1024, "Optional: Argon2 memory parameter (in KiB) used for encrypted private key passphrase")
|
||||||
cf.argonParallelism = cf.set.Uint("argon-parallelism", 4, "Optional: Argon2 parallelism parameter used for encrypted private key passphrase")
|
cf.argonParallelism = cf.set.Uint("argon-parallelism", 4, "Optional: Argon2 parallelism parameter used for encrypted private key passphrase")
|
||||||
cf.argonIterations = cf.set.Uint("argon-iterations", defaultArgonIterations, "Optional: Argon2 iterations parameter used for encrypted private key passphrase")
|
cf.argonIterations = cf.set.Uint("argon-iterations", 1, "Optional: Argon2 iterations parameter used for encrypted private key passphrase")
|
||||||
cf.encryption = cf.set.Bool("encrypt", false, "Optional: prompt for passphrase and write out-key in an encrypted format")
|
cf.encryption = cf.set.Bool("encrypt", false, "Optional: prompt for passphrase and write out-key in an encrypted format")
|
||||||
cf.curve = cf.set.String("curve", defaultCurve(), "EdDSA/ECDSA Curve (25519, P256)")
|
cf.curve = cf.set.String("curve", "25519", "EdDSA/ECDSA Curve (25519, P256)")
|
||||||
cf.p11url = p11Flag(cf.set)
|
cf.p11url = p11Flag(cf.set)
|
||||||
|
|
||||||
cf.ips = cf.set.String("ips", "", "Deprecated, see -networks")
|
cf.ips = cf.set.String("ips", "", "Deprecated, see -networks")
|
||||||
@@ -268,9 +244,6 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
} else {
|
} else {
|
||||||
switch *cf.curve {
|
switch *cf.curve {
|
||||||
case "25519", "X25519", "Curve25519", "CURVE25519":
|
case "25519", "X25519", "Curve25519", "CURVE25519":
|
||||||
if fips140.Enforced() {
|
|
||||||
return errors.New("use of Curve25519 is not allowed in FIPS 140-only mode")
|
|
||||||
}
|
|
||||||
curve = cert.Curve_CURVE25519
|
curve = cert.Curve_CURVE25519
|
||||||
pub, rawPriv, err = ed25519.GenerateKey(rand.Reader)
|
pub, rawPriv, err = ed25519.GenerateKey(rand.Reader)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -7,9 +7,7 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"encoding/pem"
|
"encoding/pem"
|
||||||
"errors"
|
"errors"
|
||||||
"math/bits"
|
|
||||||
"os"
|
"os"
|
||||||
"strconv"
|
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
@@ -24,18 +22,6 @@ func Test_caSummary(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func Test_caHelp(t *testing.T) {
|
func Test_caHelp(t *testing.T) {
|
||||||
var (
|
|
||||||
defaultArgonMemory string
|
|
||||||
defaultArgonIterations string
|
|
||||||
)
|
|
||||||
if bits.UintSize == 32 {
|
|
||||||
defaultArgonMemory = strconv.Itoa(64 * 1024)
|
|
||||||
defaultArgonIterations = strconv.Itoa(3)
|
|
||||||
} else {
|
|
||||||
defaultArgonMemory = strconv.Itoa(2 * 1024 * 1024)
|
|
||||||
defaultArgonIterations = strconv.Itoa(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
ob := &bytes.Buffer{}
|
ob := &bytes.Buffer{}
|
||||||
caHelp(ob)
|
caHelp(ob)
|
||||||
assert.Equal(
|
assert.Equal(
|
||||||
@@ -43,9 +29,9 @@ func Test_caHelp(t *testing.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"+
|
" 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 "+defaultArgonIterations+")\n"+
|
" \tOptional: Argon2 iterations parameter used for encrypted private key passphrase (default 1)\n"+
|
||||||
" -argon-memory uint\n"+
|
" -argon-memory uint\n"+
|
||||||
" \tOptional: Argon2 memory parameter (in KiB) used for encrypted private key passphrase (default "+defaultArgonMemory+")\n"+
|
" \tOptional: Argon2 memory parameter (in KiB) used for encrypted private key passphrase (default 2097152)\n"+
|
||||||
" -argon-parallelism uint\n"+
|
" -argon-parallelism uint\n"+
|
||||||
" \tOptional: Argon2 parallelism parameter used for encrypted private key passphrase (default 4)\n"+
|
" \tOptional: Argon2 parallelism parameter used for encrypted private key passphrase (default 4)\n"+
|
||||||
" -curve string\n"+
|
" -curve string\n"+
|
||||||
@@ -202,16 +188,10 @@ func Test_ca(t *testing.T) {
|
|||||||
k, _ := pem.Decode(rb)
|
k, _ := pem.Decode(rb)
|
||||||
ned, err := cert.UnmarshalNebulaEncryptedData(k.Bytes)
|
ned, err := cert.UnmarshalNebulaEncryptedData(k.Bytes)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
// we won't know salt in advance, so just check start of string
|
||||||
if bits.UintSize == 32 {
|
assert.Equal(t, uint32(2*1024*1024), ned.EncryptionMetadata.Argon2Parameters.Memory)
|
||||||
assert.Equal(t, uint32(64*1024), ned.EncryptionMetadata.Argon2Parameters.Memory)
|
|
||||||
assert.Equal(t, uint32(3), ned.EncryptionMetadata.Argon2Parameters.Iterations)
|
|
||||||
} else {
|
|
||||||
assert.Equal(t, uint32(2*1024*1024), ned.EncryptionMetadata.Argon2Parameters.Memory)
|
|
||||||
assert.Equal(t, uint32(1), ned.EncryptionMetadata.Argon2Parameters.Iterations)
|
|
||||||
}
|
|
||||||
|
|
||||||
assert.Equal(t, uint8(4), ned.EncryptionMetadata.Argon2Parameters.Parallelism)
|
assert.Equal(t, uint8(4), ned.EncryptionMetadata.Argon2Parameters.Parallelism)
|
||||||
|
assert.Equal(t, uint32(1), ned.EncryptionMetadata.Argon2Parameters.Iterations)
|
||||||
|
|
||||||
// verify the key is valid and decrypt-able
|
// verify the key is valid and decrypt-able
|
||||||
var curve cert.Curve
|
var curve cert.Curve
|
||||||
|
|||||||
@@ -1,5 +0,0 @@
|
|||||||
//go:build fips140enforce
|
|
||||||
|
|
||||||
//go:debug fips140=only
|
|
||||||
|
|
||||||
package main
|
|
||||||
@@ -1,8 +1,6 @@
|
|||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/fips140"
|
|
||||||
"errors"
|
|
||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
@@ -26,7 +24,7 @@ func newKeygenFlags() *keygenFlags {
|
|||||||
cf.set.Usage = func() {}
|
cf.set.Usage = func() {}
|
||||||
cf.outPubPath = cf.set.String("out-pub", "", "Required: path to write the public key to")
|
cf.outPubPath = cf.set.String("out-pub", "", "Required: path to write the public key to")
|
||||||
cf.outKeyPath = cf.set.String("out-key", "", "Required: path to write the private key to")
|
cf.outKeyPath = cf.set.String("out-key", "", "Required: path to write the private key to")
|
||||||
cf.curve = cf.set.String("curve", defaultCurve(), "ECDH Curve (25519, P256)")
|
cf.curve = cf.set.String("curve", "25519", "ECDH Curve (25519, P256)")
|
||||||
cf.p11url = p11Flag(cf.set)
|
cf.p11url = p11Flag(cf.set)
|
||||||
return &cf
|
return &cf
|
||||||
}
|
}
|
||||||
@@ -63,9 +61,6 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
} else {
|
} else {
|
||||||
switch *cf.curve {
|
switch *cf.curve {
|
||||||
case "25519", "X25519", "Curve25519", "CURVE25519":
|
case "25519", "X25519", "Curve25519", "CURVE25519":
|
||||||
if fips140.Enforced() {
|
|
||||||
return errors.New("use of Curve25519 is not allowed in FIPS 140-only mode")
|
|
||||||
}
|
|
||||||
pub, rawPriv = x25519Keypair()
|
pub, rawPriv = x25519Keypair()
|
||||||
curve = cert.Curve_CURVE25519
|
curve = cert.Curve_CURVE25519
|
||||||
case "P256":
|
case "P256":
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ package main
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/ecdh"
|
"crypto/ecdh"
|
||||||
"crypto/fips140"
|
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"errors"
|
"errors"
|
||||||
"flag"
|
"flag"
|
||||||
@@ -269,10 +268,6 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
}(p11Client)
|
}(p11Client)
|
||||||
}
|
}
|
||||||
|
|
||||||
if fips140.Enforced() && curve == cert.Curve_CURVE25519 {
|
|
||||||
return errors.New("use of Curve25519 is not allowed in FIPS 140-only mode")
|
|
||||||
}
|
|
||||||
|
|
||||||
if *sf.inPubPath != "" {
|
if *sf.inPubPath != "" {
|
||||||
var pubCurve cert.Curve
|
var pubCurve cert.Curve
|
||||||
rawPub, err := readInput("in-pub", *sf.inPubPath, &claims)
|
rawPub, err := readInput("in-pub", *sf.inPubPath, &claims)
|
||||||
|
|||||||
@@ -1,5 +0,0 @@
|
|||||||
//go:build fips140enforce
|
|
||||||
|
|
||||||
//go:debug fips140=only
|
|
||||||
|
|
||||||
package main
|
|
||||||
@@ -1,5 +0,0 @@
|
|||||||
//go:build fips140enforce
|
|
||||||
|
|
||||||
//go:debug fips140=only
|
|
||||||
|
|
||||||
package main
|
|
||||||
+8
-34
@@ -105,18 +105,11 @@ func (cm *connectionManager) getInactivityTimeout() time.Duration {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) In(h *HostInfo) {
|
func (cm *connectionManager) In(h *HostInfo) {
|
||||||
h.markIn()
|
h.in.Store(true)
|
||||||
}
|
}
|
||||||
|
|
||||||
// OutNoRebind records outbound traffic without consuming the rebind epoch, for relayed sends: the direct path
|
func (cm *connectionManager) Out(h *HostInfo) {
|
||||||
// to the relay consumes the edge, the via send must not.
|
h.out.Store(true)
|
||||||
func (cm *connectionManager) OutNoRebind(h *HostInfo) {
|
|
||||||
h.markOutOnly()
|
|
||||||
}
|
|
||||||
|
|
||||||
// Out records outbound traffic and reports whether we rebound since this tunnel last sent
|
|
||||||
func (cm *connectionManager) Out(h *HostInfo) bool {
|
|
||||||
return h.markOut(cm.intf.rebindEpoch.Load())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) RelayUsed(localIndex uint32) {
|
func (cm *connectionManager) RelayUsed(localIndex uint32) {
|
||||||
@@ -135,7 +128,8 @@ func (cm *connectionManager) RelayUsed(localIndex uint32) {
|
|||||||
// getAndResetTrafficCheck returns if there was any inbound or outbound traffic within the last tick and
|
// getAndResetTrafficCheck returns if there was any inbound or outbound traffic within the last tick and
|
||||||
// resets the state for this local index
|
// resets the state for this local index
|
||||||
func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time) (bool, bool) {
|
func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time) (bool, bool) {
|
||||||
in, out := h.takeTraffic()
|
in := h.in.Swap(false)
|
||||||
|
out := h.out.Swap(false)
|
||||||
if in || out {
|
if in || out {
|
||||||
h.lastUsed = now
|
h.lastUsed = now
|
||||||
}
|
}
|
||||||
@@ -329,12 +323,6 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
return closeTunnel, hostinfo, nil
|
return closeTunnel, hostinfo, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if hostinfo.ConnectionState != nil && hostinfo.ConnectionState.messageCounter.Load() >= RejectAfterMessages {
|
|
||||||
// Send path can't encrypt a CloseTunnel notify, so just delete locally; the peer recovers via recv_error.
|
|
||||||
hostinfo.logger(cm.l).Error("Dropping tunnel, message counter is exhausted")
|
|
||||||
return deleteTunnel, hostinfo, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
primary := cm.hostMap.Hosts[hostinfo.vpnAddrs[0]]
|
primary := cm.hostMap.Hosts[hostinfo.vpnAddrs[0]]
|
||||||
mainHostInfo := true
|
mainHostInfo := true
|
||||||
if primary != nil && primary != hostinfo {
|
if primary != nil && primary != hostinfo {
|
||||||
@@ -352,7 +340,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
"tunnelCheck", m{"state": "alive", "method": "passive"},
|
"tunnelCheck", m{"state": "alive", "method": "passive"},
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
hostinfo.setPendingDeletion(false)
|
hostinfo.pendingDeletion.Store(false)
|
||||||
|
|
||||||
if mainHostInfo {
|
if mainHostInfo {
|
||||||
decision = tryRehandshake
|
decision = tryRehandshake
|
||||||
@@ -375,7 +363,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
return decision, hostinfo, primary
|
return decision, hostinfo, primary
|
||||||
}
|
}
|
||||||
|
|
||||||
if hostinfo.isPendingDeletion() {
|
if hostinfo.pendingDeletion.Load() {
|
||||||
// We have already sent a test packet and nothing was returned, this hostinfo is dead
|
// We have already sent a test packet and nothing was returned, this hostinfo is dead
|
||||||
hostinfo.logger(cm.l).Info("Tunnel status",
|
hostinfo.logger(cm.l).Info("Tunnel status",
|
||||||
"tunnelCheck", m{"state": "dead", "method": "active"},
|
"tunnelCheck", m{"state": "dead", "method": "active"},
|
||||||
@@ -426,7 +414,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
hostinfo.setPendingDeletion(true)
|
hostinfo.pendingDeletion.Store(true)
|
||||||
cm.trafficTimer.Add(hostinfo.localIndexId, cm.pendingDeletionInterval)
|
cm.trafficTimer.Add(hostinfo.localIndexId, cm.pendingDeletionInterval)
|
||||||
return decision, hostinfo, nil
|
return decision, hostinfo, nil
|
||||||
}
|
}
|
||||||
@@ -460,11 +448,6 @@ func (cm *connectionManager) shouldSwapPrimary(current *HostInfo) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
if current.ConnectionState.messageCounter.Load() >= RehandshakeAfterMessages {
|
|
||||||
// This tunnel is being rolled for counter exhaustion, never swap back onto its spent key.
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
crt := cm.intf.pki.getCertState().getCertificate(current.ConnectionState.myCert.Version())
|
crt := cm.intf.pki.getCertState().getCertificate(current.ConnectionState.myCert.Version())
|
||||||
if crt == nil {
|
if crt == nil {
|
||||||
//my cert was reloaded away. We should definitely swap from this tunnel
|
//my cert was reloaded away. We should definitely swap from this tunnel
|
||||||
@@ -561,15 +544,6 @@ func (cm *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
|||||||
"reason", "current cert version < pki.initiatingVersion",
|
"reason", "current cert version < pki.initiatingVersion",
|
||||||
)
|
)
|
||||||
|
|
||||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if hostinfo.ConnectionState.messageCounter.Load() >= RehandshakeAfterMessages {
|
|
||||||
cm.l.Info("Re-handshaking with remote",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"reason", "message counter rehandshake threshold reached",
|
|
||||||
)
|
|
||||||
|
|
||||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
+36
-109
@@ -86,25 +86,25 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
// We saw traffic out to vpnIp
|
// We saw traffic out to vpnIp
|
||||||
nc.Out(hostinfo)
|
nc.Out(hostinfo)
|
||||||
nc.In(hostinfo)
|
nc.In(hostinfo)
|
||||||
assert.False(t, hostinfo.isPendingDeletion())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
assert.True(t, hostinfo.sentSinceCheck())
|
assert.True(t, hostinfo.out.Load())
|
||||||
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.True(t, hostinfo.in.Load())
|
||||||
|
|
||||||
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
|
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||||
assert.False(t, hostinfo.isPendingDeletion())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.sentSinceCheck())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.False(t, hostinfo.in.Load())
|
||||||
|
|
||||||
// Do another traffic check tick, this host should be pending deletion now
|
// Do another traffic check tick, this host should be pending deletion now
|
||||||
nc.Out(hostinfo)
|
nc.Out(hostinfo)
|
||||||
assert.True(t, hostinfo.sentSinceCheck())
|
assert.True(t, hostinfo.out.Load())
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||||
assert.True(t, hostinfo.isPendingDeletion())
|
assert.True(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.sentSinceCheck())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.False(t, hostinfo.in.Load())
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||||
|
|
||||||
@@ -168,110 +168,37 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
// We saw traffic out to vpnIp
|
// We saw traffic out to vpnIp
|
||||||
nc.Out(hostinfo)
|
nc.Out(hostinfo)
|
||||||
nc.In(hostinfo)
|
nc.In(hostinfo)
|
||||||
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.True(t, hostinfo.in.Load())
|
||||||
assert.True(t, hostinfo.sentSinceCheck())
|
assert.True(t, hostinfo.out.Load())
|
||||||
assert.False(t, hostinfo.isPendingDeletion())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
|
|
||||||
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
|
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||||
assert.False(t, hostinfo.isPendingDeletion())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.sentSinceCheck())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.False(t, hostinfo.in.Load())
|
||||||
|
|
||||||
// Do another traffic check tick, this host should be pending deletion now
|
// Do another traffic check tick, this host should be pending deletion now
|
||||||
nc.Out(hostinfo)
|
nc.Out(hostinfo)
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||||
assert.True(t, hostinfo.isPendingDeletion())
|
assert.True(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.sentSinceCheck())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.False(t, hostinfo.in.Load())
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||||
|
|
||||||
// We saw traffic, should no longer be pending deletion
|
// We saw traffic, should no longer be pending deletion
|
||||||
nc.In(hostinfo)
|
nc.In(hostinfo)
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||||
assert.False(t, hostinfo.isPendingDeletion())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.sentSinceCheck())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.False(t, hostinfo.in.Load())
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||||
}
|
}
|
||||||
|
|
||||||
func Test_NewConnectionManager_CounterLimits(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
localrange := netip.MustParsePrefix("10.1.1.1/24")
|
|
||||||
vpnIp := netip.MustParseAddr("172.1.1.2")
|
|
||||||
preferredRanges := []netip.Prefix{localrange}
|
|
||||||
|
|
||||||
// Very incomplete mock objects
|
|
||||||
hostMap := newHostMap(l)
|
|
||||||
hostMap.preferredRanges.Store(&preferredRanges)
|
|
||||||
|
|
||||||
cs := &CertState{
|
|
||||||
initiatingVersion: cert.Version1,
|
|
||||||
privateKey: []byte{},
|
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
|
||||||
v1Credential: nil,
|
|
||||||
}
|
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
|
||||||
ifce := &Interface{
|
|
||||||
hostMap: hostMap,
|
|
||||||
inside: &overlaytest.NoopTun{},
|
|
||||||
outside: &udp.NoopConn{},
|
|
||||||
firewall: &Firewall{},
|
|
||||||
lightHouse: lh,
|
|
||||||
pki: &PKI{},
|
|
||||||
myVpnAddrs: []netip.Addr{netip.MustParseAddr("172.1.1.1")}, // sorts below vpnIp so shouldSwapPrimary can proceed
|
|
||||||
handshakeManager: NewHandshakeManager(l, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
|
||||||
l: l,
|
|
||||||
}
|
|
||||||
ifce.pki.cs.Store(cs)
|
|
||||||
|
|
||||||
conf := config.NewC(test.NewLogger())
|
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
|
||||||
nc.intf = ifce
|
|
||||||
|
|
||||||
hostinfo := &HostInfo{
|
|
||||||
vpnAddrs: []netip.Addr{vpnIp},
|
|
||||||
localIndexId: 1099,
|
|
||||||
remoteIndexId: 9901,
|
|
||||||
}
|
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
|
||||||
}
|
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
|
||||||
|
|
||||||
// Below the rehandshake threshold, no handshake is started
|
|
||||||
hostinfo.ConnectionState.messageCounter.Store(RehandshakeAfterMessages - 1)
|
|
||||||
nc.tryRehandshake(hostinfo)
|
|
||||||
assert.Nil(t, ifce.handshakeManager.QueryVpnAddr(vpnIp))
|
|
||||||
|
|
||||||
// A tunnel on its current cert would normally swap to primary
|
|
||||||
assert.True(t, nc.shouldSwapPrimary(hostinfo))
|
|
||||||
|
|
||||||
// At the rehandshake threshold, a new handshake is started
|
|
||||||
hostinfo.ConnectionState.messageCounter.Store(RehandshakeAfterMessages)
|
|
||||||
nc.tryRehandshake(hostinfo)
|
|
||||||
assert.NotNil(t, ifce.handshakeManager.QueryVpnAddr(vpnIp))
|
|
||||||
|
|
||||||
// An exhausted tunnel being rolled must never swap back to primary onto its spent key
|
|
||||||
assert.False(t, nc.shouldSwapPrimary(hostinfo))
|
|
||||||
|
|
||||||
// Still below the reject limit, the tunnel stays up
|
|
||||||
nc.In(hostinfo)
|
|
||||||
decision, _, _ := nc.makeTrafficDecision(hostinfo.localIndexId, time.Now())
|
|
||||||
assert.Equal(t, tryRehandshake, decision)
|
|
||||||
|
|
||||||
// At the reject limit, the tunnel is deleted locally without a doomed CloseTunnel notify
|
|
||||||
hostinfo.ConnectionState.messageCounter.Store(RejectAfterMessages)
|
|
||||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, time.Now())
|
|
||||||
assert.Equal(t, deleteTunnel, decision)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
localrange := netip.MustParsePrefix("10.1.1.1/24")
|
localrange := netip.MustParsePrefix("10.1.1.1/24")
|
||||||
@@ -326,31 +253,31 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
// Do a traffic check tick, in and out should be cleared but should not be pending deletion
|
// Do a traffic check tick, in and out should be cleared but should not be pending deletion
|
||||||
nc.Out(hostinfo)
|
nc.Out(hostinfo)
|
||||||
nc.In(hostinfo)
|
nc.In(hostinfo)
|
||||||
assert.True(t, hostinfo.sentSinceCheck())
|
assert.True(t, hostinfo.out.Load())
|
||||||
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.True(t, hostinfo.in.Load())
|
||||||
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
decision, _, _ := nc.makeTrafficDecision(hostinfo.localIndexId, now)
|
decision, _, _ := nc.makeTrafficDecision(hostinfo.localIndexId, now)
|
||||||
assert.Equal(t, tryRehandshake, decision)
|
assert.Equal(t, tryRehandshake, decision)
|
||||||
assert.Equal(t, now, hostinfo.lastUsed)
|
assert.Equal(t, now, hostinfo.lastUsed)
|
||||||
assert.False(t, hostinfo.isPendingDeletion())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.sentSinceCheck())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.False(t, hostinfo.in.Load())
|
||||||
|
|
||||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*5))
|
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*5))
|
||||||
assert.Equal(t, doNothing, decision)
|
assert.Equal(t, doNothing, decision)
|
||||||
assert.Equal(t, now, hostinfo.lastUsed)
|
assert.Equal(t, now, hostinfo.lastUsed)
|
||||||
assert.False(t, hostinfo.isPendingDeletion())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.sentSinceCheck())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.False(t, hostinfo.in.Load())
|
||||||
|
|
||||||
// Do another traffic check tick, should still not be pending deletion
|
// Do another traffic check tick, should still not be pending deletion
|
||||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*10))
|
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*10))
|
||||||
assert.Equal(t, doNothing, decision)
|
assert.Equal(t, doNothing, decision)
|
||||||
assert.Equal(t, now, hostinfo.lastUsed)
|
assert.Equal(t, now, hostinfo.lastUsed)
|
||||||
assert.False(t, hostinfo.isPendingDeletion())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.sentSinceCheck())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.False(t, hostinfo.in.Load())
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||||
|
|
||||||
@@ -358,9 +285,9 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Minute*10))
|
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Minute*10))
|
||||||
assert.Equal(t, closeTunnel, decision)
|
assert.Equal(t, closeTunnel, decision)
|
||||||
assert.Equal(t, now, hostinfo.lastUsed)
|
assert.Equal(t, now, hostinfo.lastUsed)
|
||||||
assert.False(t, hostinfo.isPendingDeletion())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.sentSinceCheck())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
assert.False(t, hostinfo.in.Load())
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||||
}
|
}
|
||||||
|
|||||||
+8
-48
@@ -2,7 +2,6 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
@@ -13,26 +12,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const ReplayWindow = 1024
|
||||||
ReplayWindow = 8192
|
|
||||||
|
|
||||||
// RehandshakeAfterMessages rolls keys inside the AES-GCM data-volume margin (~2^-36 advantage at 64KB frames).
|
|
||||||
RehandshakeAfterMessages = uint64(1) << 34
|
|
||||||
|
|
||||||
// RejectAfterMessages is the nonce ceiling enforced by noiseutil; a tunnel here is deleted locally, not notified.
|
|
||||||
RejectAfterMessages = noiseutil.RejectAfterMessages
|
|
||||||
)
|
|
||||||
|
|
||||||
// RehandshakeAfterMessages must stay below RejectAfterMessages so tunnels roll before the hard send stop.
|
|
||||||
const _ = RejectAfterMessages - RehandshakeAfterMessages
|
|
||||||
|
|
||||||
// sessionEpoch hands out a receiver-local ordinal to every ConnectionState at creation. The RX
|
|
||||||
// staging sort (overlay/batch) orders packets by (epoch, message counter). A re-handshake never
|
|
||||||
// rekeys an existing tunnel; it brings up a new hostinfo and ConnectionState with a counter space
|
|
||||||
// starting near zero, while the old tunnel keeps decrypting until torn down. During that cutover
|
|
||||||
// one flush batch can hold packets from both tunnels, and the epoch keeps the old tunnel's
|
|
||||||
// packets sorted first.
|
|
||||||
var sessionEpoch atomic.Uint64
|
|
||||||
|
|
||||||
type ConnectionState struct {
|
type ConnectionState struct {
|
||||||
eKey noiseutil.CipherState
|
eKey noiseutil.CipherState
|
||||||
@@ -44,20 +24,13 @@ type ConnectionState struct {
|
|||||||
window *Bits
|
window *Bits
|
||||||
decryptLock sync.Mutex
|
decryptLock sync.Mutex
|
||||||
writeLock sync.Mutex
|
writeLock sync.Mutex
|
||||||
// epoch is this session's sessionEpoch ordinal. Immutable after creation.
|
|
||||||
epoch uint64
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// newConnectionStateFromResult builds a fully-populated ConnectionState from a
|
// newConnectionStateFromResult builds a fully-populated ConnectionState from a
|
||||||
// completed handshake.Result. It seeds messageCounter and the replay window so
|
// completed handshake.Result. It seeds messageCounter and the replay window so
|
||||||
// that the post-handshake message indices already used on the wire don't count
|
// that the post-handshake message indices already used on the wire don't count
|
||||||
// as missed traffic in the data plane.
|
// as missed traffic in the data plane.
|
||||||
func newConnectionStateFromResult(r *handshake.Result) (*ConnectionState, error) {
|
func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
|
||||||
// Refuse a MessageIndex too big for the replay window: it can only be a bug, and would spin the seed loop below.
|
|
||||||
if r.MessageIndex >= ReplayWindow {
|
|
||||||
return nil, fmt.Errorf("handshake message index %d exceeds replay window", r.MessageIndex)
|
|
||||||
}
|
|
||||||
|
|
||||||
ci := &ConnectionState{
|
ci := &ConnectionState{
|
||||||
myCert: r.MyCert,
|
myCert: r.MyCert,
|
||||||
initiator: r.Initiator,
|
initiator: r.Initiator,
|
||||||
@@ -65,13 +38,12 @@ func newConnectionStateFromResult(r *handshake.Result) (*ConnectionState, error)
|
|||||||
eKey: noiseutil.NewCipherState(r.EKey, r.Cipher),
|
eKey: noiseutil.NewCipherState(r.EKey, r.Cipher),
|
||||||
dKey: noiseutil.NewCipherState(r.DKey, r.Cipher),
|
dKey: noiseutil.NewCipherState(r.DKey, r.Cipher),
|
||||||
window: NewBits(ReplayWindow),
|
window: NewBits(ReplayWindow),
|
||||||
epoch: sessionEpoch.Add(1),
|
|
||||||
}
|
}
|
||||||
ci.messageCounter.Add(r.MessageIndex)
|
ci.messageCounter.Add(r.MessageIndex)
|
||||||
for i := uint64(1); i <= r.MessageIndex; i++ {
|
for i := uint64(1); i <= r.MessageIndex; i++ {
|
||||||
ci.window.Update(nil, i)
|
ci.window.Update(nil, i)
|
||||||
}
|
}
|
||||||
return ci, nil
|
return ci
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
||||||
@@ -82,21 +54,12 @@ func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// NextMessageCounter reserves the next 1-based counter; RejectAfterMessages is the first we refuse, pinned to not wrap.
|
|
||||||
func (cs *ConnectionState) NextMessageCounter() (uint64, bool) {
|
|
||||||
c := cs.messageCounter.Add(1)
|
|
||||||
if c >= RejectAfterMessages {
|
|
||||||
cs.messageCounter.Store(RejectAfterMessages)
|
|
||||||
return c, false
|
|
||||||
}
|
|
||||||
return c, true
|
|
||||||
}
|
|
||||||
|
|
||||||
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, packet []byte, nb []byte) ([]byte, error) {
|
func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out []byte, packet []byte, nb []byte) ([]byte, error) {
|
||||||
|
var err error
|
||||||
cs.decryptLock.Lock()
|
cs.decryptLock.Lock()
|
||||||
result := cs.window.Check(l, messageCounter)
|
result := cs.window.Check(l, messageCounter)
|
||||||
cs.decryptLock.Unlock()
|
cs.decryptLock.Unlock()
|
||||||
@@ -104,7 +67,7 @@ func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, packet
|
|||||||
return nil, ErrAlreadySeen
|
return nil, ErrAlreadySeen
|
||||||
}
|
}
|
||||||
|
|
||||||
out, err := cs.dKey.DecryptDanger(packet[header.Len:header.Len], packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
out, err = cs.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -118,6 +81,7 @@ func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, packet
|
|||||||
return out, nil
|
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 {
|
func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) error {
|
||||||
cs.decryptLock.Lock()
|
cs.decryptLock.Lock()
|
||||||
result := cs.window.Check(l, messageCounter)
|
result := cs.window.Check(l, messageCounter)
|
||||||
@@ -126,11 +90,6 @@ func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, pa
|
|||||||
return ErrAlreadySeen
|
return ErrAlreadySeen
|
||||||
}
|
}
|
||||||
|
|
||||||
// The entire body is sent as AD, not encrypted.
|
|
||||||
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
|
||||||
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
|
|
||||||
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
|
|
||||||
// which will gracefully fail in the DecryptDanger call.
|
|
||||||
signedPayload := packet[:len(packet)-cs.dKey.Overhead()]
|
signedPayload := packet[:len(packet)-cs.dKey.Overhead()]
|
||||||
signatureValue := packet[len(packet)-cs.dKey.Overhead():]
|
signatureValue := packet[len(packet)-cs.dKey.Overhead():]
|
||||||
_, err := cs.dKey.DecryptDanger(nil, signedPayload, signatureValue, messageCounter, nb)
|
_, err := cs.dKey.DecryptDanger(nil, signedPayload, signatureValue, messageCounter, nb)
|
||||||
@@ -144,5 +103,6 @@ func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, pa
|
|||||||
if !result {
|
if !result {
|
||||||
return ErrAlreadySeen
|
return ErrAlreadySeen
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,13 +6,10 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
"github.com/flynn/noise"
|
||||||
"github.com/rcrowley/go-metrics"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
ct "github.com/slackhq/nebula/cert_test"
|
ct "github.com/slackhq/nebula/cert_test"
|
||||||
"github.com/slackhq/nebula/handshake"
|
"github.com/slackhq/nebula/handshake"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"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"
|
||||||
)
|
)
|
||||||
@@ -82,77 +79,11 @@ func runTestHandshake(t *testing.T) (initR, respR *handshake.Result) {
|
|||||||
return initR, respR
|
return initR, respR
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConnectionState_NextMessageCounter(t *testing.T) {
|
|
||||||
cs := &ConnectionState{}
|
|
||||||
cs.messageCounter.Store(RejectAfterMessages - 2)
|
|
||||||
|
|
||||||
c, ok := cs.NextMessageCounter()
|
|
||||||
assert.True(t, ok)
|
|
||||||
assert.Equal(t, RejectAfterMessages-1, c)
|
|
||||||
|
|
||||||
// Hitting the limit refuses and pins the counter there
|
|
||||||
c, ok = cs.NextMessageCounter()
|
|
||||||
assert.False(t, ok)
|
|
||||||
assert.Equal(t, RejectAfterMessages, c)
|
|
||||||
assert.Equal(t, RejectAfterMessages, cs.messageCounter.Load())
|
|
||||||
|
|
||||||
// Continued send attempts stay refused and the counter never wraps
|
|
||||||
for i := 0; i < 10; i++ {
|
|
||||||
_, ok = cs.NextMessageCounter()
|
|
||||||
assert.False(t, ok)
|
|
||||||
}
|
|
||||||
assert.Equal(t, RejectAfterMessages, cs.messageCounter.Load())
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestSendNoMetricsDropsExhausted drives the send path to the exhausted drop; metric and out flag prove it.
|
|
||||||
func TestSendNoMetricsDropsExhausted(t *testing.T) {
|
|
||||||
initR, _ := runTestHandshake(t)
|
|
||||||
ci, err := newConnectionStateFromResult(initR)
|
|
||||||
require.NoError(t, err)
|
|
||||||
ci.messageCounter.Store(RejectAfterMessages - 1)
|
|
||||||
|
|
||||||
f := &Interface{l: test.NewLogger(), messageMetrics: &MessageMetrics{txExhausted: metrics.NewCounter()}}
|
|
||||||
hostinfo := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("10.0.0.1")}, ConnectionState: ci}
|
|
||||||
|
|
||||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, []byte{}, make([]byte, 12), make([]byte, mtu), 0)
|
|
||||||
|
|
||||||
// The crossing send is refused: it records an exhaustion drop and never reaches connectionManager.Out.
|
|
||||||
assert.Equal(t, int64(1), f.messageMetrics.txExhausted.Count())
|
|
||||||
assert.False(t, hostinfo.sentSinceCheck())
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestSendNoMetricsCloseTunnelKeepsRebindEpoch pins that a closing tunnel does not consume a rebind, a later
|
|
||||||
// packet on a re-established tunnel still needs that edge to trigger the far-side punch.
|
|
||||||
func TestSendNoMetricsCloseTunnelKeepsRebindEpoch(t *testing.T) {
|
|
||||||
initR, _ := runTestHandshake(t)
|
|
||||||
ci, err := newConnectionStateFromResult(initR)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
f := &Interface{
|
|
||||||
l: test.NewLogger(),
|
|
||||||
messageMetrics: &MessageMetrics{txExhausted: metrics.NewCounter()},
|
|
||||||
writers: []udp.Conn{udp.NoopConn{}},
|
|
||||||
connectionManager: &connectionManager{},
|
|
||||||
}
|
|
||||||
hostinfo := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("10.0.0.1")}, ConnectionState: ci}
|
|
||||||
|
|
||||||
// Tunnel is on epoch 0, then we rebind.
|
|
||||||
hostinfo.markOut(0)
|
|
||||||
f.rebindEpoch.Add(1)
|
|
||||||
|
|
||||||
remote := netip.MustParseAddrPort("10.0.0.2:4242")
|
|
||||||
f.sendNoMetrics(header.CloseTunnel, 0, ci, hostinfo, remote, []byte{}, make([]byte, 12), make([]byte, mtu), 0)
|
|
||||||
|
|
||||||
// markOut at the new epoch still reports the move, so the edge was preserved.
|
|
||||||
assert.True(t, hostinfo.markOut(1), "a CloseTunnel send must not consume the rebind epoch")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewConnectionStateFromResult(t *testing.T) {
|
func TestNewConnectionStateFromResult(t *testing.T) {
|
||||||
initR, respR := runTestHandshake(t)
|
initR, respR := runTestHandshake(t)
|
||||||
|
|
||||||
t.Run("initiator", func(t *testing.T) {
|
t.Run("initiator", func(t *testing.T) {
|
||||||
ci, err := newConnectionStateFromResult(initR)
|
ci := newConnectionStateFromResult(initR)
|
||||||
require.NoError(t, err)
|
|
||||||
assert.True(t, ci.initiator)
|
assert.True(t, ci.initiator)
|
||||||
assert.Equal(t, initR.MyCert, ci.myCert)
|
assert.Equal(t, initR.MyCert, ci.myCert)
|
||||||
assert.Equal(t, initR.RemoteCert, ci.peerCert)
|
assert.Equal(t, initR.RemoteCert, ci.peerCert)
|
||||||
@@ -171,17 +102,8 @@ func TestNewConnectionStateFromResult(t *testing.T) {
|
|||||||
assert.True(t, ci.window.Check(nil, 3), "counter 3 must not be pre-seeded")
|
assert.True(t, ci.window.Check(nil, 3), "counter 3 must not be pre-seeded")
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("message index too large is refused", func(t *testing.T) {
|
|
||||||
bad := *initR
|
|
||||||
bad.MessageIndex = ReplayWindow
|
|
||||||
ci, err := newConnectionStateFromResult(&bad)
|
|
||||||
require.Error(t, err)
|
|
||||||
assert.Nil(t, ci)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("responder", func(t *testing.T) {
|
t.Run("responder", func(t *testing.T) {
|
||||||
ci, err := newConnectionStateFromResult(respR)
|
ci := newConnectionStateFromResult(respR)
|
||||||
require.NoError(t, err)
|
|
||||||
assert.False(t, ci.initiator)
|
assert.False(t, ci.initiator)
|
||||||
assert.Equal(t, respR.MyCert, ci.myCert)
|
assert.Equal(t, respR.MyCert, ci.myCert)
|
||||||
assert.Equal(t, respR.RemoteCert, ci.peerCert)
|
assert.Equal(t, respR.RemoteCert, ci.peerCert)
|
||||||
|
|||||||
+2
-2
@@ -115,7 +115,7 @@ func (c *Control) Start() error {
|
|||||||
c.lighthouseStart()
|
c.lighthouseStart()
|
||||||
}
|
}
|
||||||
|
|
||||||
c.f.triggerShutdown = func() { go c.Stop() }
|
c.f.triggerShutdown = c.Stop
|
||||||
|
|
||||||
// Start reading packets.
|
// Start reading packets.
|
||||||
c.f.run()
|
c.f.run()
|
||||||
@@ -212,7 +212,7 @@ func (c *Control) RebindUDPServer() {
|
|||||||
c.f.lightHouse.SendUpdate()
|
c.f.lightHouse.SendUpdate()
|
||||||
|
|
||||||
// Let the main interface know that we rebound so that underlying tunnels know to trigger punches from their remotes
|
// Let the main interface know that we rebound so that underlying tunnels know to trigger punches from their remotes
|
||||||
c.f.rebindEpoch.Add(1)
|
c.f.rebindCount++
|
||||||
}
|
}
|
||||||
|
|
||||||
// ListHostmapHosts returns details about the actual or pending (handshaking) hostmap by vpn ip
|
// ListHostmapHosts returns details about the actual or pending (handshaking) hostmap by vpn ip
|
||||||
|
|||||||
+23
-40
@@ -11,8 +11,6 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/overlay/batch"
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
@@ -32,9 +30,9 @@ func newFakeDevice() *fakeDevice {
|
|||||||
|
|
||||||
// Read blocks until Close like a real tun with no traffic, then reports EOF
|
// Read blocks until Close like a real tun with no traffic, then reports EOF
|
||||||
// the same way a closed device does
|
// the same way a closed device does
|
||||||
func (d *fakeDevice) Read() ([]tio.Packet, error) {
|
func (d *fakeDevice) Read(p []byte) (int, error) {
|
||||||
<-d.closedCh
|
<-d.closedCh
|
||||||
return nil, io.EOF
|
return 0, io.EOF
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *fakeDevice) Write(p []byte) (int, error) { return len(p), nil }
|
func (d *fakeDevice) Write(p []byte) (int, error) { return len(p), nil }
|
||||||
@@ -51,8 +49,10 @@ func (d *fakeDevice) Activate() error { return nil }
|
|||||||
func (d *fakeDevice) Networks() []netip.Prefix { return nil }
|
func (d *fakeDevice) Networks() []netip.Prefix { return nil }
|
||||||
func (d *fakeDevice) Name() string { return "fake" }
|
func (d *fakeDevice) Name() string { return "fake" }
|
||||||
func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil }
|
func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil }
|
||||||
|
func (d *fakeDevice) SupportsMultiqueue() bool { return false }
|
||||||
func (d *fakeDevice) Queues(int) ([]tio.Queue, error) { return []tio.Queue{d}, nil }
|
func (d *fakeDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
|
return nil, errors.New("unsupported")
|
||||||
|
}
|
||||||
|
|
||||||
// newReadyControl hand-builds the minimum Control that Main would have
|
// newReadyControl hand-builds the minimum Control that Main would have
|
||||||
// produced right before Start, including the construction token NewInterface
|
// produced right before Start, including the construction token NewInterface
|
||||||
@@ -78,7 +78,7 @@ func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
|
|||||||
inside: dev,
|
inside: dev,
|
||||||
outside: conn,
|
outside: conn,
|
||||||
writers: []udp.Conn{conn},
|
writers: []udp.Conn{conn},
|
||||||
batchers: make([]*batch.MultiCoalescer, 1),
|
readers: make([]io.ReadWriteCloser, 1),
|
||||||
routines: 1,
|
routines: 1,
|
||||||
hostMap: newHostMap(l),
|
hostMap: newHostMap(l),
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
@@ -109,8 +109,7 @@ func TestControl_StopBeforeStart(t *testing.T) {
|
|||||||
require.NoError(t, c.Wait())
|
require.NoError(t, c.Wait())
|
||||||
|
|
||||||
// A stopped control can never be started
|
// A stopped control can never be started
|
||||||
err := c.Start()
|
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||||
require.ErrorIs(t, err, ErrAlreadyStopped)
|
|
||||||
|
|
||||||
// A second Stop is a harmless no-op
|
// A second Stop is a harmless no-op
|
||||||
c.Stop()
|
c.Stop()
|
||||||
@@ -144,29 +143,19 @@ type fakeConn struct {
|
|||||||
rebinds int
|
rebinds int
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
|
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
|
||||||
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||||
func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil }
|
func (c *fakeConn) ListenOut(_ udp.EncReader) error { return nil }
|
||||||
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
||||||
func (c *fakeConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) {
|
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
||||||
return len(bufs), nil
|
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
||||||
}
|
func (c *fakeConn) Close() error { c.closed = true; 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 {
|
type multiqueueDevice struct {
|
||||||
*fakeDevice
|
*fakeDevice
|
||||||
}
|
}
|
||||||
|
|
||||||
// Queues claims multiqueue support but fails to open the second queue,
|
func (d *multiqueueDevice) SupportsMultiqueue() bool { return true }
|
||||||
// exercising the activation error path.
|
|
||||||
func (d *multiqueueDevice) Queues(n int) ([]tio.Queue, error) {
|
|
||||||
if n > 1 {
|
|
||||||
return nil, errors.New("second queue failed to open")
|
|
||||||
}
|
|
||||||
return d.fakeDevice.Queues(n)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
||||||
dev := &multiqueueDevice{fakeDevice: newFakeDevice()}
|
dev := &multiqueueDevice{fakeDevice: newFakeDevice()}
|
||||||
@@ -177,7 +166,7 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
|||||||
inside: dev,
|
inside: dev,
|
||||||
outside: conn,
|
outside: conn,
|
||||||
writers: []udp.Conn{conn},
|
writers: []udp.Conn{conn},
|
||||||
batchers: make([]*batch.MultiCoalescer, 2),
|
readers: make([]io.ReadWriteCloser, 2),
|
||||||
routines: 2,
|
routines: 2,
|
||||||
l: test.NewLogger(),
|
l: test.NewLogger(),
|
||||||
}
|
}
|
||||||
@@ -192,8 +181,7 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// The second reader fails to open, everything must be released
|
// The second reader fails to open, everything must be released
|
||||||
err := c.Start()
|
require.Error(t, c.Start())
|
||||||
require.Error(t, err)
|
|
||||||
assert.Equal(t, StateStopped, c.State())
|
assert.Equal(t, StateStopped, c.State())
|
||||||
assert.True(t, dev.closed, "the tun device should have been closed")
|
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||||
assert.True(t, conn.closed, "the udp socket should have been closed")
|
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||||
@@ -263,18 +251,15 @@ func TestControl_ConcurrentStopAndStart(t *testing.T) {
|
|||||||
// panic and Wait must observe the final state
|
// panic and Wait must observe the final state
|
||||||
require.NoError(t, c.Wait())
|
require.NoError(t, c.Wait())
|
||||||
assert.Equal(t, StateStopped, c.State())
|
assert.Equal(t, StateStopped, c.State())
|
||||||
err := c.Start()
|
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||||
require.ErrorIs(t, err, ErrAlreadyStopped)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestControl_StartStopLifecycle(t *testing.T) {
|
func TestControl_StartStopLifecycle(t *testing.T) {
|
||||||
c, dev, conn := newReadyControl(t)
|
c, dev, conn := newReadyControl(t)
|
||||||
|
|
||||||
err := c.Start()
|
require.NoError(t, c.Start())
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, StateStarted, c.State())
|
assert.Equal(t, StateStarted, c.State())
|
||||||
err = c.Start()
|
require.ErrorIs(t, c.Start(), ErrAlreadyStarted)
|
||||||
require.ErrorIs(t, err, ErrAlreadyStarted)
|
|
||||||
|
|
||||||
// Stop must unpark the reader blocked in the device and release everything
|
// Stop must unpark the reader blocked in the device and release everything
|
||||||
c.Stop()
|
c.Stop()
|
||||||
@@ -285,8 +270,7 @@ func TestControl_StartStopLifecycle(t *testing.T) {
|
|||||||
|
|
||||||
// The reader drained off a closed device, that is not a fatal error
|
// The reader drained off a closed device, that is not a fatal error
|
||||||
require.NoError(t, c.Wait())
|
require.NoError(t, c.Wait())
|
||||||
err = c.Start()
|
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||||
require.ErrorIs(t, err, ErrAlreadyStopped)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestControl_RebindIsGatedByState(t *testing.T) {
|
func TestControl_RebindIsGatedByState(t *testing.T) {
|
||||||
@@ -296,8 +280,7 @@ func TestControl_RebindIsGatedByState(t *testing.T) {
|
|||||||
c.RebindUDPServer()
|
c.RebindUDPServer()
|
||||||
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
|
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
|
||||||
|
|
||||||
err := c.Start()
|
require.NoError(t, c.Start())
|
||||||
require.NoError(t, err)
|
|
||||||
c.RebindUDPServer()
|
c.RebindUDPServer()
|
||||||
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
|
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
|
||||||
|
|
||||||
|
|||||||
@@ -123,16 +123,6 @@ func (c *Control) SetLocalAddrsFn(fn func(*LocalAllowList) []netip.Addr) {
|
|||||||
c.f.lightHouse.localAddrsFn = fn
|
c.f.lightHouse.localAddrsFn = fn
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetRebindEpochFor returns the rebind epoch a tunnel last sent under, so a test can tell whether a send
|
|
||||||
// consumed the epoch edge without having to infer it from lighthouse traffic.
|
|
||||||
func (c *Control) GetRebindEpochFor(vpnAddr netip.Addr) (uint32, bool) {
|
|
||||||
h := c.f.hostMap.QueryVpnAddr(vpnAddr)
|
|
||||||
if h == nil {
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
return h.state.Load() >> stateEpochShift, true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
|
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
|
||||||
hostinfo := c.f.handshakeManager.QueryVpnAddr(vpnIp)
|
hostinfo := c.f.handshakeManager.QueryVpnAddr(vpnIp)
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
|
|||||||
@@ -1,187 +0,0 @@
|
|||||||
// Package cpupick chooses which CPUs the tun reader threads pin to when the
|
|
||||||
// operator has not chosen for us (tun.cpu_affinity). The stock spread —
|
|
||||||
// allowed[i] for routine i — has two failure modes this package exists to fix:
|
|
||||||
//
|
|
||||||
// - every co-located nebula starts its spread at allowed[0], so N instances
|
|
||||||
// on one box stack their readers onto the same cores, and allowed[0] is
|
|
||||||
// usually CPU 0, the core housekeeping and default IRQ affinity already
|
|
||||||
// favor;
|
|
||||||
// - on heterogeneous CPUs (ARM big.LITTLE, Intel P/E hybrids, AMD compact
|
|
||||||
// cores) low IDs are not necessarily fast cores, and pinning an encrypt
|
|
||||||
// thread to an efficiency core caps that queue's throughput.
|
|
||||||
//
|
|
||||||
// Default instead returns a preference-ordered pin list: the allowed set
|
|
||||||
// filtered to performance cores (when the platform distinguishes them and
|
|
||||||
// enough remain for every routine), confined to a single NUMA node and spread
|
|
||||||
// across distinct physical cores when the topology permits, CPU 0's physical
|
|
||||||
// core demoted to last resort, and the order rotated by a stable per-instance
|
|
||||||
// key so co-located instances spread instead of stacking.
|
|
||||||
package cpupick
|
|
||||||
|
|
||||||
import (
|
|
||||||
"log/slog"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/util"
|
|
||||||
)
|
|
||||||
|
|
||||||
// topology is the slice of machine layout arrange consults: the NUMA node
|
|
||||||
// and the physical core behind each candidate CPU, plus which core CPU 0
|
|
||||||
// lives on (zeroCore, -1 when unknown — tracked separately because CPU 0's
|
|
||||||
// SMT sibling deserves demotion even when CPU 0 itself isn't a candidate).
|
|
||||||
// Probed from sysfs on Linux; flatTopology stands in when the platform can't
|
|
||||||
// say, which turns every topology rule into a no-op rather than a wrong
|
|
||||||
// answer.
|
|
||||||
type topology struct {
|
|
||||||
nodeOf map[int]int
|
|
||||||
coreOf map[int]int
|
|
||||||
zeroCore int
|
|
||||||
}
|
|
||||||
|
|
||||||
// flatTopology places every CPU on node 0 and on a physical core of its own.
|
|
||||||
func flatTopology(cpus []int) topology {
|
|
||||||
t := topology{
|
|
||||||
nodeOf: make(map[int]int, len(cpus)),
|
|
||||||
coreOf: make(map[int]int, len(cpus)),
|
|
||||||
zeroCore: -1,
|
|
||||||
}
|
|
||||||
for i, c := range cpus {
|
|
||||||
t.nodeOf[c] = 0
|
|
||||||
t.coreOf[c] = i
|
|
||||||
if c == 0 {
|
|
||||||
t.zeroCore = i
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return t
|
|
||||||
}
|
|
||||||
|
|
||||||
// Default computes the pin order for `routines` tun readers. key is any
|
|
||||||
// stable per-instance value; the bound UDP port is ideal — distinct across
|
|
||||||
// co-located instances, stable across restarts so benchmark runs stay
|
|
||||||
// comparable. Returns nil when there is nothing useful to say (no affinity
|
|
||||||
// support on this platform, lookup failure); callers keep their existing
|
|
||||||
// fallback spread.
|
|
||||||
func Default(routines int, key uint64, l *slog.Logger) []int {
|
|
||||||
allowed, err := util.AllowedCPUs()
|
|
||||||
if err != nil || len(allowed) == 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
perf, signal := perfCPUs(allowed)
|
|
||||||
cands := pickCandidates(allowed, perf, routines)
|
|
||||||
if len(cands) == 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
if len(perf) < routines {
|
|
||||||
signal = ""
|
|
||||||
}
|
|
||||||
cpus := arrange(cands, readTopology(cands), routines, splitmix64(key))
|
|
||||||
if l != nil {
|
|
||||||
l.Info("chose default pin CPUs for tun readers",
|
|
||||||
"cpus", cpus[:min(routines, len(cpus))],
|
|
||||||
"perfSignal", signal)
|
|
||||||
}
|
|
||||||
return cpus
|
|
||||||
}
|
|
||||||
|
|
||||||
// pickCandidates applies the enough-for-everyone guard: a perf filter that
|
|
||||||
// leaves fewer candidates than routines is discarded — giving every reader
|
|
||||||
// its own (possibly slow) core beats stacking two readers on a fast one.
|
|
||||||
func pickCandidates(allowed, perf []int, routines int) []int {
|
|
||||||
if len(perf) < routines {
|
|
||||||
return allowed
|
|
||||||
}
|
|
||||||
return perf
|
|
||||||
}
|
|
||||||
|
|
||||||
// arrange turns the candidate set into the final pin order:
|
|
||||||
//
|
|
||||||
// 1. NUMA: when at least one node holds enough candidates for every
|
|
||||||
// routine, confine to one such node, chosen by the instance hash. The
|
|
||||||
// readers share hostmap and cipher state, so splitting one instance
|
|
||||||
// across nodes taxes every packet — and co-located instances that hash
|
|
||||||
// to different nodes stop competing entirely. When no node is big
|
|
||||||
// enough, span nodes rather than stack readers.
|
|
||||||
// 2. Rotate the preferred candidates by the hash so instances spread.
|
|
||||||
// 3. SMT: emit one thread per physical core before any of their siblings —
|
|
||||||
// two encrypt threads on one core split its execution units. Siblings
|
|
||||||
// still follow for the routines > cores case.
|
|
||||||
// 4. CPU 0's whole physical core goes last: housekeeping and default IRQ
|
|
||||||
// noise on CPU 0 bleeds into its SMT sibling too. Within that tail the
|
|
||||||
// sibling precedes CPU 0 itself, which only catches the bleed-through.
|
|
||||||
//
|
|
||||||
// The rotation happens before the SMT pass so each instance's one-per-core
|
|
||||||
// walk also starts at a different core, and CPU 0's core is excluded from
|
|
||||||
// the rotation so no hash value can put it back at the front.
|
|
||||||
func arrange(cands []int, topo topology, routines int, h uint64) []int {
|
|
||||||
byNode := map[int][]int{}
|
|
||||||
var nodes []int
|
|
||||||
for _, c := range cands {
|
|
||||||
n := topo.nodeOf[c]
|
|
||||||
if _, ok := byNode[n]; !ok {
|
|
||||||
nodes = append(nodes, n)
|
|
||||||
}
|
|
||||||
byNode[n] = append(byNode[n], c)
|
|
||||||
}
|
|
||||||
var eligible []int
|
|
||||||
for _, n := range nodes {
|
|
||||||
if len(byNode[n]) >= routines {
|
|
||||||
eligible = append(eligible, n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if len(eligible) > 0 {
|
|
||||||
cands = byNode[eligible[int(h%uint64(len(eligible)))]]
|
|
||||||
}
|
|
||||||
|
|
||||||
// Split off CPU 0's core: its siblings tail the list, CPU 0 tails them.
|
|
||||||
preferred := make([]int, 0, len(cands))
|
|
||||||
var zeroTail []int
|
|
||||||
hasZero := false
|
|
||||||
for _, c := range cands {
|
|
||||||
switch {
|
|
||||||
case c == 0:
|
|
||||||
hasZero = true
|
|
||||||
case topo.zeroCore >= 0 && topo.coreOf[c] == topo.zeroCore:
|
|
||||||
zeroTail = append(zeroTail, c)
|
|
||||||
default:
|
|
||||||
preferred = append(preferred, c)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if hasZero {
|
|
||||||
zeroTail = append(zeroTail, 0)
|
|
||||||
}
|
|
||||||
if len(preferred) == 0 {
|
|
||||||
return zeroTail // CPU 0's core is all we have
|
|
||||||
}
|
|
||||||
|
|
||||||
// The node pick consumed the low hash bits; rotate by the high ones so
|
|
||||||
// the two choices stay independent.
|
|
||||||
off := int((h >> 32) % uint64(len(preferred)))
|
|
||||||
rot := make([]int, 0, len(preferred))
|
|
||||||
rot = append(rot, preferred[off:]...)
|
|
||||||
rot = append(rot, preferred[:off]...)
|
|
||||||
|
|
||||||
seenCore := make(map[int]bool, len(rot))
|
|
||||||
out := make([]int, 0, len(cands))
|
|
||||||
var siblings []int
|
|
||||||
for _, c := range rot {
|
|
||||||
g := topo.coreOf[c]
|
|
||||||
if seenCore[g] {
|
|
||||||
siblings = append(siblings, c)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
seenCore[g] = true
|
|
||||||
out = append(out, c)
|
|
||||||
}
|
|
||||||
out = append(out, siblings...)
|
|
||||||
out = append(out, zeroTail...)
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
// splitmix64 decorrelates instance keys before the selection modulos: ports
|
|
||||||
// on one box often share spacing (4242/4243, or round steps like +1000) that
|
|
||||||
// raw key%len arithmetic would fold onto the same offset.
|
|
||||||
func splitmix64(x uint64) uint64 {
|
|
||||||
x += 0x9e3779b97f4a7c15
|
|
||||||
x = (x ^ (x >> 30)) * 0xbf58476d1ce4e5b9
|
|
||||||
x = (x ^ (x >> 27)) * 0x94d049bb133111eb
|
|
||||||
return x ^ (x >> 31)
|
|
||||||
}
|
|
||||||
@@ -1,171 +0,0 @@
|
|||||||
package cpupick
|
|
||||||
|
|
||||||
import (
|
|
||||||
"slices"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
// pairTopo builds a topology where consecutive candidate pairs are SMT
|
|
||||||
// siblings: (cpus[0],cpus[1]) share a core, (cpus[2],cpus[3]) the next, ...
|
|
||||||
// All CPUs land on node 0.
|
|
||||||
func pairTopo(cpus []int) topology {
|
|
||||||
t := topology{
|
|
||||||
nodeOf: make(map[int]int, len(cpus)),
|
|
||||||
coreOf: make(map[int]int, len(cpus)),
|
|
||||||
zeroCore: -1,
|
|
||||||
}
|
|
||||||
for i, c := range cpus {
|
|
||||||
t.nodeOf[c] = 0
|
|
||||||
t.coreOf[c] = i / 2
|
|
||||||
if c == 0 {
|
|
||||||
t.zeroCore = i / 2
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return t
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestArrangeDemotesZeroForEveryKey(t *testing.T) {
|
|
||||||
candidates := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
|
||||||
for key := range uint64(64) {
|
|
||||||
got := arrange(candidates, flatTopology(candidates), 4, splitmix64(key))
|
|
||||||
if len(got) != len(candidates) {
|
|
||||||
t.Fatalf("key %d: len=%d want %d", key, len(got), len(candidates))
|
|
||||||
}
|
|
||||||
if got[0] == 0 {
|
|
||||||
t.Errorf("key %d: CPU 0 at the front: %v", key, got)
|
|
||||||
}
|
|
||||||
if got[len(got)-1] != 0 {
|
|
||||||
t.Errorf("key %d: CPU 0 not demoted to last: %v", key, got)
|
|
||||||
}
|
|
||||||
sorted := slices.Clone(got)
|
|
||||||
slices.Sort(sorted)
|
|
||||||
if !slices.Equal(sorted, candidates) {
|
|
||||||
t.Errorf("key %d: not a permutation: %v", key, got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestArrangeDemotesZeroSiblings(t *testing.T) {
|
|
||||||
// Pairs (0,1),(2,3),(4,5),(6,7): CPU 0's core — 0 and its sibling 1 —
|
|
||||||
// must tail the list, sibling ahead of 0 itself.
|
|
||||||
candidates := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
|
||||||
for key := range uint64(64) {
|
|
||||||
got := arrange(candidates, pairTopo(candidates), 2, splitmix64(key))
|
|
||||||
n := len(got)
|
|
||||||
if got[n-1] != 0 || got[n-2] != 1 {
|
|
||||||
t.Fatalf("key %d: tail = %v, want [... 1 0]", key, got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestArrangeZeroSiblingWithoutZero(t *testing.T) {
|
|
||||||
// CPU 0 excluded (cpuset) but its sibling 1 remains: the sibling still
|
|
||||||
// tails the list when the topology knows which core CPU 0 lives on.
|
|
||||||
candidates := []int{1, 2, 3, 4, 5}
|
|
||||||
topo := pairTopo([]int{0, 1, 2, 3, 4, 5})
|
|
||||||
got := arrange(candidates, topo, 2, splitmix64(7))
|
|
||||||
if got[len(got)-1] != 1 {
|
|
||||||
t.Errorf("CPU 0's sibling not demoted: %v", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestArrangeRotatesByKey(t *testing.T) {
|
|
||||||
candidates := []int{1, 2, 3, 4, 5, 6, 7, 8}
|
|
||||||
seen := map[int]bool{}
|
|
||||||
for key := range uint64(64) {
|
|
||||||
seen[arrange(candidates, flatTopology(candidates), 4, splitmix64(key))[0]] = true
|
|
||||||
}
|
|
||||||
// 64 hashed keys over 8 slots must hit more than one starting CPU, or
|
|
||||||
// co-located instances would all stack again.
|
|
||||||
if len(seen) < 2 {
|
|
||||||
t.Errorf("rotation never varied across keys: %v", seen)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestArrangeStableForSameKey(t *testing.T) {
|
|
||||||
candidates := []int{0, 2, 4, 6}
|
|
||||||
topo := flatTopology(candidates)
|
|
||||||
a := arrange(candidates, topo, 2, splitmix64(4242))
|
|
||||||
b := arrange(candidates, topo, 2, splitmix64(4242))
|
|
||||||
if !slices.Equal(a, b) {
|
|
||||||
t.Errorf("same key ordered differently: %v vs %v", a, b)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestArrangeZeroOnly(t *testing.T) {
|
|
||||||
if got := arrange([]int{0}, flatTopology([]int{0}), 1, splitmix64(7)); !slices.Equal(got, []int{0}) {
|
|
||||||
t.Errorf("sole CPU 0 must survive: %v", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestArrangeSMTSiblingsLast(t *testing.T) {
|
|
||||||
// Pairs (1,2),(3,4),(5,6),(7,8): the first four picks must cover four
|
|
||||||
// distinct physical cores before any sibling repeats.
|
|
||||||
candidates := []int{1, 2, 3, 4, 5, 6, 7, 8}
|
|
||||||
topo := pairTopo(candidates)
|
|
||||||
for key := range uint64(16) {
|
|
||||||
got := arrange(candidates, topo, 4, splitmix64(key))
|
|
||||||
seen := map[int]bool{}
|
|
||||||
for _, c := range got[:4] {
|
|
||||||
g := topo.coreOf[c]
|
|
||||||
if seen[g] {
|
|
||||||
t.Fatalf("key %d: sibling before all cores covered: %v", key, got)
|
|
||||||
}
|
|
||||||
seen[g] = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestArrangeNUMAConfinesToOneNode(t *testing.T) {
|
|
||||||
// Two nodes of four; both fit routines=3, so the result must sit
|
|
||||||
// entirely inside one of them, and the hash must pick both across keys.
|
|
||||||
candidates := []int{1, 2, 3, 4, 10, 11, 12, 13}
|
|
||||||
topo := flatTopology(candidates)
|
|
||||||
for _, c := range []int{10, 11, 12, 13} {
|
|
||||||
topo.nodeOf[c] = 1
|
|
||||||
}
|
|
||||||
nodesSeen := map[int]bool{}
|
|
||||||
for key := range uint64(32) {
|
|
||||||
got := arrange(candidates, topo, 3, splitmix64(key))
|
|
||||||
if len(got) != 4 {
|
|
||||||
t.Fatalf("key %d: not confined to one node: %v", key, got)
|
|
||||||
}
|
|
||||||
n := topo.nodeOf[got[0]]
|
|
||||||
for _, c := range got {
|
|
||||||
if topo.nodeOf[c] != n {
|
|
||||||
t.Fatalf("key %d: spans nodes: %v", key, got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
nodesSeen[n] = true
|
|
||||||
}
|
|
||||||
if len(nodesSeen) != 2 {
|
|
||||||
t.Errorf("hash never spread instances across nodes: %v", nodesSeen)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestArrangeNUMASpansWhenNoNodeFits(t *testing.T) {
|
|
||||||
candidates := []int{1, 2, 3, 4, 10, 11, 12, 13}
|
|
||||||
topo := flatTopology(candidates)
|
|
||||||
for _, c := range []int{10, 11, 12, 13} {
|
|
||||||
topo.nodeOf[c] = 1
|
|
||||||
}
|
|
||||||
got := arrange(candidates, topo, 6, splitmix64(1))
|
|
||||||
if len(got) != len(candidates) {
|
|
||||||
t.Errorf("undersized nodes must span, got %v", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPickCandidates(t *testing.T) {
|
|
||||||
allowed := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
|
||||||
perf := []int{4, 5}
|
|
||||||
|
|
||||||
// Enough perf cores for every routine: only they are used.
|
|
||||||
if got := pickCandidates(allowed, perf, 2); !slices.Equal(got, perf) {
|
|
||||||
t.Errorf("perf filter not applied: %v", got)
|
|
||||||
}
|
|
||||||
// Perf filter too small for the routine count: discarded, everyone
|
|
||||||
// gets their own core from the full allowed set.
|
|
||||||
if got := pickCandidates(allowed, perf, 4); !slices.Equal(got, allowed) {
|
|
||||||
t.Errorf("undersized perf filter not discarded: %v", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,154 +0,0 @@
|
|||||||
//go:build linux
|
|
||||||
|
|
||||||
package cpupick
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
)
|
|
||||||
|
|
||||||
// capacityKeepPct is the cpu_capacity admission threshold, relative to the
|
|
||||||
// fastest allowed core. LITTLE cores are normalized to ~250-400 of the big
|
|
||||||
// core's 1024 while mid cores sit at ~75%+, so half of max separates little
|
|
||||||
// from the rest without splitting prime from mid on three-tier parts.
|
|
||||||
const capacityKeepPct = 50
|
|
||||||
|
|
||||||
// freqKeepPct is the cpuinfo_max_freq admission threshold. Favored-core
|
|
||||||
// turbo skew is 2-4% and ARM mid-vs-prime ~12%, while E-cores, LITTLE
|
|
||||||
// cores, and AMD compact cores all sit >= 20% below their siblings' max.
|
|
||||||
const freqKeepPct = 85
|
|
||||||
|
|
||||||
// perfCPUs partitions allowed into the subset that are "performance" cores,
|
|
||||||
// consulting (in order of authority):
|
|
||||||
//
|
|
||||||
// 1. cpu_capacity — arch_topology's normalized per-CPU capacity, exposed on
|
|
||||||
// arm/arm64/riscv; the scheduler's own view of big vs LITTLE.
|
|
||||||
// 2. /sys/devices/cpu_core/cpus — the Intel hybrid P-core PMU mask, present
|
|
||||||
// only on P/E parts (x86 has no cpu_capacity) and naming P cores outright.
|
|
||||||
// 3. cpuinfo_max_freq — the cross-vendor fallback; catches AMD compact
|
|
||||||
// cores, which neither of the above covers.
|
|
||||||
//
|
|
||||||
// Returns allowed unchanged (signal "") when nothing distinguishes the
|
|
||||||
// cores: homogeneous parts, VMs without cpufreq, sysfs unavailable.
|
|
||||||
func perfCPUs(allowed []int) ([]int, string) {
|
|
||||||
return perfCPUsFrom("/sys/devices/system/cpu", "/sys/devices/cpu_core/cpus", allowed)
|
|
||||||
}
|
|
||||||
|
|
||||||
func perfCPUsFrom(cpuDir, intelCoreMask string, allowed []int) ([]int, string) {
|
|
||||||
if cpus, ok := byPerCPUValue(cpuDir, "cpu_capacity", allowed, capacityKeepPct); ok {
|
|
||||||
return cpus, "cpu_capacity"
|
|
||||||
}
|
|
||||||
if cpus, ok := byIntelCoreMask(intelCoreMask, allowed); ok {
|
|
||||||
return cpus, "intel_core_pmu"
|
|
||||||
}
|
|
||||||
if cpus, ok := byPerCPUValue(cpuDir, "cpufreq/cpuinfo_max_freq", allowed, freqKeepPct); ok {
|
|
||||||
return cpus, "max_freq"
|
|
||||||
}
|
|
||||||
return allowed, ""
|
|
||||||
}
|
|
||||||
|
|
||||||
// byPerCPUValue keeps the allowed CPUs whose per-CPU sysfs value is at least
|
|
||||||
// keepPct percent of the maximum across allowed. Inconclusive (ok=false)
|
|
||||||
// when any CPU is missing the file or when every value is equal.
|
|
||||||
func byPerCPUValue(cpuDir, file string, allowed []int, keepPct int) ([]int, bool) {
|
|
||||||
vals := make([]int, len(allowed))
|
|
||||||
minV, maxV := 0, 0
|
|
||||||
for i, cpu := range allowed {
|
|
||||||
v, err := readIntFile(filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), file))
|
|
||||||
if err != nil {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
vals[i] = v
|
|
||||||
if i == 0 || v < minV {
|
|
||||||
minV = v
|
|
||||||
}
|
|
||||||
if v > maxV {
|
|
||||||
maxV = v
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if minV == maxV {
|
|
||||||
return nil, false // homogeneous by this signal; try the next one
|
|
||||||
}
|
|
||||||
keep := make([]int, 0, len(allowed))
|
|
||||||
for i, cpu := range allowed {
|
|
||||||
if vals[i]*100 >= maxV*keepPct {
|
|
||||||
keep = append(keep, cpu)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return keep, true
|
|
||||||
}
|
|
||||||
|
|
||||||
// byIntelCoreMask keeps the allowed CPUs named by the hybrid P-core PMU
|
|
||||||
// mask. Inconclusive when the file is absent (non-hybrid x86, other arches)
|
|
||||||
// or no allowed CPU is in the mask (the process was deliberately confined
|
|
||||||
// to E-cores; nothing useful to prefer within that).
|
|
||||||
func byIntelCoreMask(maskPath string, allowed []int) ([]int, bool) {
|
|
||||||
b, err := os.ReadFile(maskPath)
|
|
||||||
if err != nil {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
set, err := parseCPUList(strings.TrimSpace(string(b)))
|
|
||||||
if err != nil || len(set) == 0 {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
pcore := make(map[int]bool, len(set))
|
|
||||||
for _, c := range set {
|
|
||||||
pcore[c] = true
|
|
||||||
}
|
|
||||||
keep := make([]int, 0, len(allowed))
|
|
||||||
for _, cpu := range allowed {
|
|
||||||
if pcore[cpu] {
|
|
||||||
keep = append(keep, cpu)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if len(keep) == 0 {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
return keep, true
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseCPUList decodes the kernel's cpulist format ("0-7,16-23", "3") into
|
|
||||||
// individual CPU IDs. Empty input yields an empty list.
|
|
||||||
func parseCPUList(s string) ([]int, error) {
|
|
||||||
if s == "" {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
var out []int
|
|
||||||
for part := range strings.SplitSeq(s, ",") {
|
|
||||||
part = strings.TrimSpace(part)
|
|
||||||
if part == "" {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
lo, hi, isRange := strings.Cut(part, "-")
|
|
||||||
a, err := strconv.Atoi(lo)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err)
|
|
||||||
}
|
|
||||||
if !isRange {
|
|
||||||
out = append(out, a)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
b, err := strconv.Atoi(hi)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err)
|
|
||||||
}
|
|
||||||
if b < a || b-a > 8192 {
|
|
||||||
return nil, fmt.Errorf("bad cpulist range %q", part)
|
|
||||||
}
|
|
||||||
for v := a; v <= b; v++ {
|
|
||||||
out = append(out, v)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func readIntFile(path string) (int, error) {
|
|
||||||
b, err := os.ReadFile(path)
|
|
||||||
if err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
return strconv.Atoi(strings.TrimSpace(string(b)))
|
|
||||||
}
|
|
||||||
@@ -1,163 +0,0 @@
|
|||||||
//go:build linux
|
|
||||||
|
|
||||||
package cpupick
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"slices"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
// fakeSysfs builds a cpuDir tree with the given per-CPU file values.
|
|
||||||
// A nil map for a file means "file absent on every CPU".
|
|
||||||
func fakeSysfs(t *testing.T, capacity, maxFreq map[int]int) string {
|
|
||||||
t.Helper()
|
|
||||||
dir := t.TempDir()
|
|
||||||
write := func(cpu int, rel string, v int) {
|
|
||||||
p := filepath.Join(dir, fmt.Sprintf("cpu%d", cpu), rel)
|
|
||||||
if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := os.WriteFile(p, fmt.Appendf(nil, "%d\n", v), 0o644); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
for cpu, v := range capacity {
|
|
||||||
write(cpu, "cpu_capacity", v)
|
|
||||||
}
|
|
||||||
for cpu, v := range maxFreq {
|
|
||||||
write(cpu, "cpufreq/cpuinfo_max_freq", v)
|
|
||||||
}
|
|
||||||
return dir
|
|
||||||
}
|
|
||||||
|
|
||||||
func writeCoreMask(t *testing.T, mask string) string {
|
|
||||||
t.Helper()
|
|
||||||
p := filepath.Join(t.TempDir(), "cpus")
|
|
||||||
if err := os.WriteFile(p, []byte(mask+"\n"), 0o644); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
return p
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPerfCPUsBigLittleCapacity(t *testing.T) {
|
|
||||||
// 4 big (1024) + 4 LITTLE (~290): capacity is authoritative on ARM.
|
|
||||||
dir := fakeSysfs(t, map[int]int{
|
|
||||||
0: 1024, 1: 1024, 2: 1024, 3: 1024,
|
|
||||||
4: 290, 5: 290, 6: 290, 7: 290,
|
|
||||||
}, nil)
|
|
||||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7})
|
|
||||||
if signal != "cpu_capacity" {
|
|
||||||
t.Fatalf("signal = %q", signal)
|
|
||||||
}
|
|
||||||
if !slices.Equal(got, []int{0, 1, 2, 3}) {
|
|
||||||
t.Errorf("got %v", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPerfCPUsThreeTierKeepsMid(t *testing.T) {
|
|
||||||
// prime (1024) + mid (~780) + little (~280): 50% keeps prime+mid.
|
|
||||||
dir := fakeSysfs(t, map[int]int{
|
|
||||||
0: 280, 1: 280, 2: 280, 3: 280,
|
|
||||||
4: 780, 5: 780, 6: 780,
|
|
||||||
7: 1024,
|
|
||||||
}, nil)
|
|
||||||
got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7})
|
|
||||||
if !slices.Equal(got, []int{4, 5, 6, 7}) {
|
|
||||||
t.Errorf("got %v", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPerfCPUsIntelHybridMask(t *testing.T) {
|
|
||||||
// No cpu_capacity on x86; the P-core PMU mask decides.
|
|
||||||
dir := fakeSysfs(t, nil, nil)
|
|
||||||
mask := writeCoreMask(t, "0-7")
|
|
||||||
got, signal := perfCPUsFrom(dir, mask, []int{0, 1, 2, 3, 8, 9, 10, 11})
|
|
||||||
if signal != "intel_core_pmu" {
|
|
||||||
t.Fatalf("signal = %q", signal)
|
|
||||||
}
|
|
||||||
if !slices.Equal(got, []int{0, 1, 2, 3}) {
|
|
||||||
t.Errorf("got %v", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPerfCPUsIntelMaskDisjointFallsThrough(t *testing.T) {
|
|
||||||
// Confined to E-cores only: the mask can't help, and equal freqs below
|
|
||||||
// mean nothing else distinguishes them either -> allowed unchanged.
|
|
||||||
dir := fakeSysfs(t, nil, map[int]int{8: 4300000, 9: 4300000})
|
|
||||||
mask := writeCoreMask(t, "0-7")
|
|
||||||
got, signal := perfCPUsFrom(dir, mask, []int{8, 9})
|
|
||||||
if signal != "" || !slices.Equal(got, []int{8, 9}) {
|
|
||||||
t.Errorf("got %v signal %q", got, signal)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPerfCPUsMaxFreqCompactCores(t *testing.T) {
|
|
||||||
// AMD-style compact cores: no capacity, no Intel mask; 3.3 vs 5.7 GHz.
|
|
||||||
dir := fakeSysfs(t, nil, map[int]int{
|
|
||||||
0: 5700000, 1: 5700000, 2: 3300000, 3: 3300000,
|
|
||||||
})
|
|
||||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3})
|
|
||||||
if signal != "max_freq" {
|
|
||||||
t.Fatalf("signal = %q", signal)
|
|
||||||
}
|
|
||||||
if !slices.Equal(got, []int{0, 1}) {
|
|
||||||
t.Errorf("got %v", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPerfCPUsFavoredCoreSkewKept(t *testing.T) {
|
|
||||||
// Turbo Boost Max favored cores run a few percent hot; they must not
|
|
||||||
// shrink the candidate set to one or two cores.
|
|
||||||
dir := fakeSysfs(t, nil, map[int]int{
|
|
||||||
0: 5800000, 1: 5700000, 2: 5700000, 3: 5600000,
|
|
||||||
})
|
|
||||||
got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3})
|
|
||||||
if !slices.Equal(got, []int{0, 1, 2, 3}) {
|
|
||||||
t.Errorf("favored-core skew filtered CPUs: %v", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPerfCPUsHomogeneousInconclusive(t *testing.T) {
|
|
||||||
dir := fakeSysfs(t, nil, map[int]int{0: 3000000, 1: 3000000})
|
|
||||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1})
|
|
||||||
if signal != "" || !slices.Equal(got, []int{0, 1}) {
|
|
||||||
t.Errorf("got %v signal %q", got, signal)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPerfCPUsNoSysfs(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2})
|
|
||||||
if signal != "" || !slices.Equal(got, []int{0, 1, 2}) {
|
|
||||||
t.Errorf("got %v signal %q", got, signal)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestParseCPUList(t *testing.T) {
|
|
||||||
cases := []struct {
|
|
||||||
in string
|
|
||||||
want []int
|
|
||||||
wantErr bool
|
|
||||||
}{
|
|
||||||
{"0-3", []int{0, 1, 2, 3}, false},
|
|
||||||
{"0-1,16-17", []int{0, 1, 16, 17}, false},
|
|
||||||
{"5", []int{5}, false},
|
|
||||||
{"", nil, false},
|
|
||||||
{"3-1", nil, true},
|
|
||||||
{"a-b", nil, true},
|
|
||||||
{"1,x", nil, true},
|
|
||||||
}
|
|
||||||
for _, c := range cases {
|
|
||||||
got, err := parseCPUList(c.in)
|
|
||||||
if (err != nil) != c.wantErr {
|
|
||||||
t.Errorf("%q: err=%v wantErr=%v", c.in, err, c.wantErr)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if !c.wantErr && !slices.Equal(got, c.want) {
|
|
||||||
t.Errorf("%q: got %v want %v", c.in, got, c.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,10 +0,0 @@
|
|||||||
//go:build !linux
|
|
||||||
|
|
||||||
package cpupick
|
|
||||||
|
|
||||||
// perfCPUs is Linux-only sysfs walking; elsewhere report "no distinction".
|
|
||||||
// Default already returns nil off-Linux (util.AllowedCPUs has no answer
|
|
||||||
// there), so this exists to keep the package compiling everywhere.
|
|
||||||
func perfCPUs(allowed []int) ([]int, string) {
|
|
||||||
return allowed, ""
|
|
||||||
}
|
|
||||||
@@ -1,118 +0,0 @@
|
|||||||
//go:build linux
|
|
||||||
|
|
||||||
package cpupick
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
)
|
|
||||||
|
|
||||||
// readTopology probes the NUMA node and physical-core layout of cpus from
|
|
||||||
// sysfs. Anything sysfs won't say degrades toward flatTopology: an unknown
|
|
||||||
// node becomes node 0, an unknown core becomes a core of its own — either
|
|
||||||
// way the corresponding arrange rule becomes a no-op instead of a wrong
|
|
||||||
// answer.
|
|
||||||
func readTopology(cpus []int) topology {
|
|
||||||
return readTopologyFrom("/sys/devices/system/node", "/sys/devices/system/cpu", cpus)
|
|
||||||
}
|
|
||||||
|
|
||||||
func readTopologyFrom(nodeDir, cpuDir string, cpus []int) topology {
|
|
||||||
coreOf, zeroCore := coreGroups(cpuDir, cpus)
|
|
||||||
return topology{
|
|
||||||
nodeOf: numaNodes(nodeDir, cpus),
|
|
||||||
coreOf: coreOf,
|
|
||||||
zeroCore: zeroCore,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// numaNodes maps each cpu to its NUMA node via
|
|
||||||
// /sys/devices/system/node/nodeN/cpulist. CPUs no node claims (or no node
|
|
||||||
// dirs at all: VMs, non-NUMA kernels) land on node 0.
|
|
||||||
func numaNodes(nodeDir string, cpus []int) map[int]int {
|
|
||||||
out := make(map[int]int, len(cpus))
|
|
||||||
for _, c := range cpus {
|
|
||||||
out[c] = 0
|
|
||||||
}
|
|
||||||
entries, err := os.ReadDir(nodeDir)
|
|
||||||
if err != nil {
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
want := make(map[int]bool, len(cpus))
|
|
||||||
for _, c := range cpus {
|
|
||||||
want[c] = true
|
|
||||||
}
|
|
||||||
for _, e := range entries {
|
|
||||||
id, ok := strings.CutPrefix(e.Name(), "node")
|
|
||||||
if !ok {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
n, err := strconv.Atoi(id)
|
|
||||||
if err != nil {
|
|
||||||
continue // has_cpu, possible, ... share the prefix
|
|
||||||
}
|
|
||||||
b, err := os.ReadFile(filepath.Join(nodeDir, e.Name(), "cpulist"))
|
|
||||||
if err != nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
list, err := parseCPUList(strings.TrimSpace(string(b)))
|
|
||||||
if err != nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
for _, c := range list {
|
|
||||||
if want[c] {
|
|
||||||
out[c] = n
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
// coreGroups maps each cpu to a dense physical-core id derived from its
|
|
||||||
// (physical_package_id, core_id) pair — core_id alone repeats across
|
|
||||||
// sockets. CPUs whose topology files are unreadable get a core of their own.
|
|
||||||
// The second return is the group id of the core CPU 0 lives on, or -1 when
|
|
||||||
// that can't be determined; CPU 0's own files are consulted even when 0 is
|
|
||||||
// not a candidate, so its SMT siblings are recognized under cpusets that
|
|
||||||
// exclude CPU 0 itself.
|
|
||||||
func coreGroups(cpuDir string, cpus []int) (map[int]int, int) {
|
|
||||||
type pkgCore struct{ pkg, core int }
|
|
||||||
pairOf := func(cpu int) (pkgCore, bool) {
|
|
||||||
topoDir := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology")
|
|
||||||
pkg, err1 := readIntFile(filepath.Join(topoDir, "physical_package_id"))
|
|
||||||
core, err2 := readIntFile(filepath.Join(topoDir, "core_id"))
|
|
||||||
if err1 != nil || err2 != nil {
|
|
||||||
return pkgCore{}, false
|
|
||||||
}
|
|
||||||
return pkgCore{pkg, core}, true
|
|
||||||
}
|
|
||||||
|
|
||||||
ids := map[pkgCore]int{}
|
|
||||||
out := make(map[int]int, len(cpus))
|
|
||||||
next := 0
|
|
||||||
for _, cpu := range cpus {
|
|
||||||
k, ok := pairOf(cpu)
|
|
||||||
if !ok {
|
|
||||||
out[cpu] = next
|
|
||||||
next++
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
id, ok := ids[k]
|
|
||||||
if !ok {
|
|
||||||
id = next
|
|
||||||
next++
|
|
||||||
ids[k] = id
|
|
||||||
}
|
|
||||||
out[cpu] = id
|
|
||||||
}
|
|
||||||
|
|
||||||
zeroCore := -1
|
|
||||||
if k, ok := pairOf(0); ok {
|
|
||||||
if id, ok := ids[k]; ok {
|
|
||||||
zeroCore = id
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return out, zeroCore
|
|
||||||
}
|
|
||||||
@@ -1,111 +0,0 @@
|
|||||||
//go:build linux
|
|
||||||
|
|
||||||
package cpupick
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
// fakeTopoSysfs builds nodeDir/cpuDir trees. nodes maps node id -> cpulist
|
|
||||||
// string; cores maps cpu -> (package, core) pair.
|
|
||||||
func fakeTopoSysfs(t *testing.T, nodes map[int]string, cores map[int][2]int) (string, string) {
|
|
||||||
t.Helper()
|
|
||||||
base := t.TempDir()
|
|
||||||
nodeDir := filepath.Join(base, "node")
|
|
||||||
cpuDir := filepath.Join(base, "cpu")
|
|
||||||
for n, list := range nodes {
|
|
||||||
d := filepath.Join(nodeDir, fmt.Sprintf("node%d", n))
|
|
||||||
if err := os.MkdirAll(d, 0o755); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := os.WriteFile(filepath.Join(d, "cpulist"), []byte(list+"\n"), 0o644); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
for cpu, pc := range cores {
|
|
||||||
d := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology")
|
|
||||||
if err := os.MkdirAll(d, 0o755); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := os.WriteFile(filepath.Join(d, "physical_package_id"), fmt.Appendf(nil, "%d\n", pc[0]), 0o644); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := os.WriteFile(filepath.Join(d, "core_id"), fmt.Appendf(nil, "%d\n", pc[1]), 0o644); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nodeDir, cpuDir
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestReadTopology(t *testing.T) {
|
|
||||||
// Two nodes; SMT pairs (0,4),(1,5) on node 0 and (2,6),(3,7) on node 1.
|
|
||||||
// core_id repeats across packages on purpose: the pair must disambiguate.
|
|
||||||
nodeDir, cpuDir := fakeTopoSysfs(t,
|
|
||||||
map[int]string{0: "0-1,4-5", 1: "2-3,6-7"},
|
|
||||||
map[int][2]int{
|
|
||||||
0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1},
|
|
||||||
2: {1, 0}, 6: {1, 0}, 3: {1, 1}, 7: {1, 1},
|
|
||||||
})
|
|
||||||
cpus := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
|
||||||
topo := readTopologyFrom(nodeDir, cpuDir, cpus)
|
|
||||||
|
|
||||||
for _, c := range []int{0, 1, 4, 5} {
|
|
||||||
if topo.nodeOf[c] != 0 {
|
|
||||||
t.Errorf("cpu %d on node %d, want 0", c, topo.nodeOf[c])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
for _, c := range []int{2, 3, 6, 7} {
|
|
||||||
if topo.nodeOf[c] != 1 {
|
|
||||||
t.Errorf("cpu %d on node %d, want 1", c, topo.nodeOf[c])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
pairs := [][2]int{{0, 4}, {1, 5}, {2, 6}, {3, 7}}
|
|
||||||
for _, p := range pairs {
|
|
||||||
if topo.coreOf[p[0]] != topo.coreOf[p[1]] {
|
|
||||||
t.Errorf("siblings %v not grouped: %d vs %d", p, topo.coreOf[p[0]], topo.coreOf[p[1]])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if topo.coreOf[0] == topo.coreOf[2] {
|
|
||||||
t.Error("cross-package cores with equal core_id must not merge")
|
|
||||||
}
|
|
||||||
if topo.zeroCore != topo.coreOf[0] {
|
|
||||||
t.Errorf("zeroCore = %d, want %d", topo.zeroCore, topo.coreOf[0])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestReadTopologyZeroCoreWithoutZeroCandidate(t *testing.T) {
|
|
||||||
// CPU 0 is not a candidate (cpuset excludes it) but its sibling 4 is:
|
|
||||||
// zeroCore must still identify their shared core.
|
|
||||||
nodeDir, cpuDir := fakeTopoSysfs(t,
|
|
||||||
map[int]string{0: "0-7"},
|
|
||||||
map[int][2]int{0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1}})
|
|
||||||
topo := readTopologyFrom(nodeDir, cpuDir, []int{1, 4, 5})
|
|
||||||
if topo.zeroCore < 0 || topo.coreOf[4] != topo.zeroCore {
|
|
||||||
t.Errorf("zeroCore = %d, coreOf[4] = %d; sibling of CPU 0 not identified", topo.zeroCore, topo.coreOf[4])
|
|
||||||
}
|
|
||||||
if topo.coreOf[1] == topo.zeroCore {
|
|
||||||
t.Error("cpu 1 wrongly grouped with CPU 0's core")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestReadTopologyMissingSysfs(t *testing.T) {
|
|
||||||
base := t.TempDir()
|
|
||||||
cpus := []int{0, 1, 2}
|
|
||||||
topo := readTopologyFrom(filepath.Join(base, "nope"), filepath.Join(base, "also-nope"), cpus)
|
|
||||||
seen := map[int]bool{}
|
|
||||||
for _, c := range cpus {
|
|
||||||
if topo.nodeOf[c] != 0 {
|
|
||||||
t.Errorf("cpu %d node = %d, want 0", c, topo.nodeOf[c])
|
|
||||||
}
|
|
||||||
if seen[topo.coreOf[c]] {
|
|
||||||
t.Errorf("cpu %d shares a fallback core group", c)
|
|
||||||
}
|
|
||||||
seen[topo.coreOf[c]] = true
|
|
||||||
}
|
|
||||||
if topo.zeroCore != -1 {
|
|
||||||
t.Errorf("zeroCore = %d, want -1 when unknown", topo.zeroCore)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,10 +0,0 @@
|
|||||||
//go:build !linux
|
|
||||||
|
|
||||||
package cpupick
|
|
||||||
|
|
||||||
// readTopology has no sysfs to consult off Linux; the flat stand-in makes
|
|
||||||
// arrange's NUMA and SMT rules no-ops. Default is already nil off Linux
|
|
||||||
// (util.AllowedCPUs has no answer there) — this keeps the package compiling.
|
|
||||||
func readTopology(cpus []int) topology {
|
|
||||||
return flatTopology(cpus)
|
|
||||||
}
|
|
||||||
+4
-2
@@ -4,13 +4,15 @@
|
|||||||
package e2e
|
package e2e
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"log/slog"
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"log/slog"
|
||||||
|
|
||||||
"dario.cat/mergo"
|
"dario.cat/mergo"
|
||||||
"github.com/google/gopacket"
|
"github.com/google/gopacket"
|
||||||
"github.com/google/gopacket/layers"
|
"github.com/google/gopacket/layers"
|
||||||
@@ -380,7 +382,7 @@ func getAddrs(ns []netip.Prefix) []netip.Addr {
|
|||||||
func NewTestLogger() *slog.Logger {
|
func NewTestLogger() *slog.Logger {
|
||||||
v := os.Getenv("TEST_LOGS")
|
v := os.Getenv("TEST_LOGS")
|
||||||
if v == "" {
|
if v == "" {
|
||||||
return slog.New(slog.DiscardHandler)
|
return slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
level := slog.LevelInfo
|
level := slog.LevelInfo
|
||||||
|
|||||||
@@ -223,60 +223,3 @@ func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
|
|||||||
lhControl.Stop()
|
lhControl.Stop()
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
// A relayed send records traffic but must not consume the rebind epoch. If it does, the next direct send to the
|
|
||||||
// relay host sees the epoch already current and never requeries, so the far side is never told to punch at our
|
|
||||||
// new address. This pins the SendVia call site, which the unit tests cannot reach.
|
|
||||||
func TestRebindRequeriesAfterRelayedSend(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{})
|
|
||||||
|
|
||||||
// No lighthouse on purpose: it would hand out a direct address for them and nothing would relay.
|
|
||||||
// Long connection manager timers so it never fires a direct test packet at the relay tunnel and bumps its
|
|
||||||
// epoch mid-test, which is the only other thing that touches that tunnel and would flake the assertion below.
|
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24",
|
|
||||||
m{"relay": m{"use_relays": true}, "timers": m{"connection_alive_interval": 3600, "pending_deletion_interval": 3600}})
|
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
|
||||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
|
||||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
|
|
||||||
r := router.NewR(t, myControl, relayControl, theirControl)
|
|
||||||
defer r.RenderFlow()
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
relayControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("establish")))
|
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
|
||||||
r.RouteFor(time.Millisecond * 500)
|
|
||||||
|
|
||||||
hi := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
|
|
||||||
require.NotNil(t, hi, "expected a tunnel to them")
|
|
||||||
require.NotEmpty(t, hi.CurrentRelaysToMe, "them must be reachable only via the relay for this test to mean anything")
|
|
||||||
// sendNoMetrics only reaches SendVia when there is no direct remote, so pin that too. Without this the test
|
|
||||||
// keeps passing while quietly sending direct and never exercising the relay path.
|
|
||||||
require.False(t, hi.CurrentRemote.IsValid(), "them must have no direct remote, otherwise SendVia is never called")
|
|
||||||
|
|
||||||
before, ok := myControl.GetRebindEpochFor(relayVpnIpNet[0].Addr())
|
|
||||||
require.True(t, ok, "expected a tunnel to the relay")
|
|
||||||
|
|
||||||
myControl.RebindUDPServer()
|
|
||||||
|
|
||||||
// Traffic to them goes through SendVia on the relay tunnel. That must record traffic without consuming the
|
|
||||||
// relay tunnel's own epoch edge, which belongs to the direct path.
|
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("relayed")))
|
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
|
||||||
|
|
||||||
after, ok := myControl.GetRebindEpochFor(relayVpnIpNet[0].Addr())
|
|
||||||
require.True(t, ok)
|
|
||||||
assert.Equal(t, before, after,
|
|
||||||
"a relayed send consumed the relay tunnel's rebind epoch, so the next direct send will not requery")
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
relayControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -131,9 +131,6 @@ listen:
|
|||||||
port: 4242
|
port: 4242
|
||||||
# Sets the max number of packets to pull from the kernel for each syscall (under systems that support recvmmsg)
|
# Sets the max number of packets to pull from the kernel for each syscall (under systems that support recvmmsg)
|
||||||
# default is 64, does not support reload
|
# default is 64, does not support reload
|
||||||
# Note: on Linux with UDP GRO (kernel 5.10+), each receive slot is sized for a full 64KiB coalesced
|
|
||||||
# superpacket, so the receive scratch is batch * 64KiB per listening socket (~4MiB per routine at the
|
|
||||||
# default of 64). Lower this to trade peak per-syscall throughput for memory on constrained hosts.
|
|
||||||
#batch: 64
|
#batch: 64
|
||||||
# Configure socket buffers for the udp side (outside), leave unset to use the system defaults. Values will be doubled by the kernel
|
# Configure socket buffers for the udp side (outside), leave unset to use the system defaults. Values will be doubled by the kernel
|
||||||
# Default is net.core.rmem_default and net.core.wmem_default (/proc/sys/net/core/rmem_default and /proc/sys/net/core/rmem_default)
|
# Default is net.core.rmem_default and net.core.wmem_default (/proc/sys/net/core/rmem_default and /proc/sys/net/core/rmem_default)
|
||||||
@@ -172,8 +169,6 @@ listen:
|
|||||||
# allowing for more precise routing decisions based on the packet tags. Default is 0 meaning no mark is set.
|
# allowing for more precise routing decisions based on the packet tags. Default is 0 meaning no mark is set.
|
||||||
# This setting is reloadable.
|
# This setting is reloadable.
|
||||||
#so_mark: 0
|
#so_mark: 0
|
||||||
# the udp_offloads setting controls if Nebula will attempt to enable GSO and GRO for its UDP socket(s). Linux only, not reloadable.
|
|
||||||
# udp_offloads: false
|
|
||||||
|
|
||||||
# Routines is the number of thread pairs to run that consume from the tun and UDP queues.
|
# Routines is the number of thread pairs to run that consume from the tun and UDP queues.
|
||||||
# Currently, this defaults to 1 which means we have 1 tun queue reader and 1
|
# Currently, this defaults to 1 which means we have 1 tun queue reader and 1
|
||||||
@@ -267,33 +262,6 @@ tun:
|
|||||||
# Default MTU for every packet, safe setting is (and the default) 1300 for internet based traffic
|
# Default MTU for every packet, safe setting is (and the default) 1300 for internet based traffic
|
||||||
mtu: 1300
|
mtu: 1300
|
||||||
|
|
||||||
# the use_offloads setting controls if Nebula will attempt to enable GSO and GRO for the tun device. Linux only, not reloadable.
|
|
||||||
#use_offloads: false
|
|
||||||
|
|
||||||
# Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's
|
|
||||||
# batched sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire
|
|
||||||
# instead of being sprayed across multiple TX rings and reordered. Not reloadable. Coerced to false if routines <= 1.
|
|
||||||
#pin_threads: true
|
|
||||||
|
|
||||||
# pin_threads_key helps the CPU-auto-selector shuffle which CPUs are chosen for pinning.
|
|
||||||
# Valid options are "pid" or "port". Use "port" if you want Nebula to choose the same cores every time, which is nice for benchmarking.
|
|
||||||
# Linux only, not reloadable.
|
|
||||||
#pin_threads_key: "pid"
|
|
||||||
|
|
||||||
# Linux only. cpu_affinity overrides which CPUs the tun reader threads pin to: a list of CPU IDs, one per routine
|
|
||||||
# (see the top-level `routines` setting). Lists shorter than `routines` are modulo-cycled across the queues; extra
|
|
||||||
# entries are ignored. IDs must be within the process's allowed CPU set, so this respects taskset / cgroup cpusets;
|
|
||||||
# a non-integer or not-allowed entry disables the override, leaving the default pin selection described below.
|
|
||||||
# Only meaningful while pin_threads is true. Not reloadable.
|
|
||||||
# When unset (or rejected), the default spread prefers performance cores on heterogeneous CPUs (ARM big.LITTLE,
|
|
||||||
# Intel P/E hybrids, AMD compact cores), keeps all readers on one NUMA node and on distinct physical cores when the
|
|
||||||
# topology allows (SMT siblings last), leaves CPU 0's physical core as a last resort, and rotates its starting
|
|
||||||
# point per instance (keyed by the bound UDP port) so co-located nebulas don't stack their readers onto the
|
|
||||||
# same cores.
|
|
||||||
#cpu_affinity:
|
|
||||||
# - 2
|
|
||||||
# - 4
|
|
||||||
|
|
||||||
# Route based MTU overrides, you have known vpn ip paths that can support larger MTUs you can increase/decrease them here
|
# Route based MTU overrides, you have known vpn ip paths that can support larger MTUs you can increase/decrease them here
|
||||||
routes:
|
routes:
|
||||||
#- mtu: 8800
|
#- mtu: 8800
|
||||||
|
|||||||
+14
-15
@@ -21,7 +21,6 @@ import (
|
|||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/iputil"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type FirewallInterface interface {
|
type FirewallInterface interface {
|
||||||
@@ -263,11 +262,11 @@ func (f *Firewall) AddRule(incoming bool, proto uint8, startPort int32, endPort
|
|||||||
}
|
}
|
||||||
|
|
||||||
switch proto {
|
switch proto {
|
||||||
case iputil.IPProtocolTCP:
|
case firewall.ProtoTCP:
|
||||||
fp = ft.TCP
|
fp = ft.TCP
|
||||||
case iputil.IPProtocolUDP:
|
case firewall.ProtoUDP:
|
||||||
fp = ft.UDP
|
fp = ft.UDP
|
||||||
case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
|
case firewall.ProtoICMP, firewall.ProtoICMPv6:
|
||||||
//ICMP traffic doesn't have ports, so we always coerce to "any", even if a value is provided
|
//ICMP traffic doesn't have ports, so we always coerce to "any", even if a value is provided
|
||||||
if startPort != firewall.PortAny {
|
if startPort != firewall.PortAny {
|
||||||
f.l.Warn("ignoring port specification for ICMP firewall rule", "startPort", startPort)
|
f.l.Warn("ignoring port specification for ICMP firewall rule", "startPort", startPort)
|
||||||
@@ -365,13 +364,13 @@ func AddFirewallRulesFromConfig(l *slog.Logger, inbound bool, c *config.C, fw Fi
|
|||||||
proto = firewall.ProtoAny
|
proto = firewall.ProtoAny
|
||||||
startPort, endPort, err = parsePort(sPort)
|
startPort, endPort, err = parsePort(sPort)
|
||||||
case "tcp":
|
case "tcp":
|
||||||
proto = iputil.IPProtocolTCP
|
proto = firewall.ProtoTCP
|
||||||
startPort, endPort, err = parsePort(sPort)
|
startPort, endPort, err = parsePort(sPort)
|
||||||
case "udp":
|
case "udp":
|
||||||
proto = iputil.IPProtocolUDP
|
proto = firewall.ProtoUDP
|
||||||
startPort, endPort, err = parsePort(sPort)
|
startPort, endPort, err = parsePort(sPort)
|
||||||
case "icmp":
|
case "icmp":
|
||||||
proto = iputil.IPProtocolICMP
|
proto = firewall.ProtoICMP
|
||||||
startPort = firewall.PortAny
|
startPort = firewall.PortAny
|
||||||
endPort = firewall.PortAny
|
endPort = firewall.PortAny
|
||||||
if sPort != "" {
|
if sPort != "" {
|
||||||
@@ -561,9 +560,9 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
|
|||||||
}
|
}
|
||||||
|
|
||||||
switch fp.Protocol {
|
switch fp.Protocol {
|
||||||
case iputil.IPProtocolTCP:
|
case firewall.ProtoTCP:
|
||||||
c.Expires = time.Now().Add(f.TCPTimeout)
|
c.Expires = time.Now().Add(f.TCPTimeout)
|
||||||
case iputil.IPProtocolUDP:
|
case firewall.ProtoUDP:
|
||||||
c.Expires = time.Now().Add(f.UDPTimeout)
|
c.Expires = time.Now().Add(f.UDPTimeout)
|
||||||
default:
|
default:
|
||||||
c.Expires = time.Now().Add(f.DefaultTimeout)
|
c.Expires = time.Now().Add(f.DefaultTimeout)
|
||||||
@@ -583,9 +582,9 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
|
|||||||
c := &conn{}
|
c := &conn{}
|
||||||
|
|
||||||
switch fp.Protocol {
|
switch fp.Protocol {
|
||||||
case iputil.IPProtocolTCP:
|
case firewall.ProtoTCP:
|
||||||
timeout = f.TCPTimeout
|
timeout = f.TCPTimeout
|
||||||
case iputil.IPProtocolUDP:
|
case firewall.ProtoUDP:
|
||||||
timeout = f.UDPTimeout
|
timeout = f.UDPTimeout
|
||||||
default:
|
default:
|
||||||
timeout = f.DefaultTimeout
|
timeout = f.DefaultTimeout
|
||||||
@@ -636,15 +635,15 @@ func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedC
|
|||||||
}
|
}
|
||||||
|
|
||||||
switch p.Protocol {
|
switch p.Protocol {
|
||||||
case iputil.IPProtocolTCP:
|
case firewall.ProtoTCP:
|
||||||
if ft.TCP.match(p, incoming, c, caPool) {
|
if ft.TCP.match(p, incoming, c, caPool) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
case iputil.IPProtocolUDP:
|
case firewall.ProtoUDP:
|
||||||
if ft.UDP.match(p, incoming, c, caPool) {
|
if ft.UDP.match(p, incoming, c, caPool) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
|
case firewall.ProtoICMP, firewall.ProtoICMPv6:
|
||||||
if ft.ICMP.match(p, incoming, c, caPool) {
|
if ft.ICMP.match(p, incoming, c, caPool) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
@@ -681,7 +680,7 @@ func (fp firewallPort) match(p firewall.Packet, incoming bool, c *cert.CachedCer
|
|||||||
}
|
}
|
||||||
|
|
||||||
// this branch is here to catch traffic from FirewallTable.Any.match and FirewallTable.ICMP.match
|
// this branch is here to catch traffic from FirewallTable.Any.match and FirewallTable.ICMP.match
|
||||||
if p.Protocol == iputil.IPProtocolICMP || p.Protocol == iputil.IPProtocolICMPv6 {
|
if p.Protocol == firewall.ProtoICMP || p.Protocol == firewall.ProtoICMPv6 {
|
||||||
// port numbers are re-used for connection tracking of ICMP,
|
// port numbers are re-used for connection tracking of ICMP,
|
||||||
// but we don't want to actually filter on them.
|
// but we don't want to actually filter on them.
|
||||||
return fp[firewall.PortAny].match(p, c, caPool)
|
return fp[firewall.PortAny].match(p, c, caPool)
|
||||||
|
|||||||
+2
-4
@@ -5,8 +5,6 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// ConntrackCache is used as a local routine cache to know if a given flow
|
// ConntrackCache is used as a local routine cache to know if a given flow
|
||||||
@@ -58,8 +56,8 @@ func (c *ConntrackCacheTicker) Get() ConntrackCache {
|
|||||||
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
||||||
c.cacheV = tick
|
c.cacheV = tick
|
||||||
if ll := len(c.cache); ll > 0 {
|
if ll := len(c.cache); ll > 0 {
|
||||||
if c.l.Enabled(context.Background(), logging.LevelTrace) {
|
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
c.l.Log(context.Background(), logging.LevelTrace, "resetting conntrack cache", "len", ll)
|
c.l.Debug("resetting conntrack cache", "len", ll)
|
||||||
}
|
}
|
||||||
c.cache = make(ConntrackCache, ll)
|
c.cache = make(ConntrackCache, ll)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
@@ -31,27 +30,27 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
|
|||||||
|
|
||||||
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 3)
|
c := newFixedTicker(t, l, 3)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|
||||||
assert.Equal(t, "level=DEBUG-4 msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
assert.Equal(t, "level=DEBUG msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewJSONLoggerWithOutput(buf, logging.LevelTrace)
|
l := test.NewJSONLoggerWithOutput(buf, slog.LevelDebug)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 2)
|
c := newFixedTicker(t, l, 2)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|
||||||
assert.JSONEq(t, `{"level":"DEBUG-4","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
assert.JSONEq(t, `{"level":"DEBUG","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) {
|
func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelInfo)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 5)
|
c := newFixedTicker(t, l, 5)
|
||||||
c.Get()
|
c.Get()
|
||||||
@@ -61,7 +60,7 @@ func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) {
|
|||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 0)
|
c := newFixedTicker(t, l, 0)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|||||||
+10
-16
@@ -4,14 +4,17 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/iputil"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type m = map[string]any
|
type m = map[string]any
|
||||||
|
|
||||||
const (
|
const (
|
||||||
ProtoAny = 0 // When we want to handle HOPOPT (0) we can change this, if ever
|
ProtoAny = 0 // When we want to handle HOPOPT (0) we can change this, if ever
|
||||||
|
ProtoTCP = 6
|
||||||
|
ProtoUDP = 17
|
||||||
|
ProtoICMP = 1
|
||||||
|
ProtoICMPv6 = 58
|
||||||
|
|
||||||
PortAny = 0 // Special value for matching `port: any`
|
PortAny = 0 // Special value for matching `port: any`
|
||||||
PortFragment = -1 // Special value for matching `port: fragment`
|
PortFragment = -1 // Special value for matching `port: fragment`
|
||||||
)
|
)
|
||||||
@@ -42,13 +45,13 @@ func (fp *Packet) Copy() *Packet {
|
|||||||
func (fp Packet) MarshalJSON() ([]byte, error) {
|
func (fp Packet) MarshalJSON() ([]byte, error) {
|
||||||
var proto string
|
var proto string
|
||||||
switch fp.Protocol {
|
switch fp.Protocol {
|
||||||
case iputil.IPProtocolTCP:
|
case ProtoTCP:
|
||||||
proto = "tcp"
|
proto = "tcp"
|
||||||
case iputil.IPProtocolICMP:
|
case ProtoICMP:
|
||||||
proto = "icmp"
|
proto = "icmp"
|
||||||
case iputil.IPProtocolICMPv6:
|
case ProtoICMPv6:
|
||||||
proto = "icmpv6"
|
proto = "icmpv6"
|
||||||
case iputil.IPProtocolUDP:
|
case ProtoUDP:
|
||||||
proto = "udp"
|
proto = "udp"
|
||||||
default:
|
default:
|
||||||
proto = fmt.Sprintf("unknown %v", fp.Protocol)
|
proto = fmt.Sprintf("unknown %v", fp.Protocol)
|
||||||
@@ -62,12 +65,3 @@ func (fp Packet) MarshalJSON() ([]byte, error) {
|
|||||||
"Fragment": fp.Fragment,
|
"Fragment": fp.Fragment,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// ParsedPacket is a Packet plus the parse byproducts the RX path reuses
|
|
||||||
type ParsedPacket struct {
|
|
||||||
Packet
|
|
||||||
IPHdrLen int
|
|
||||||
// FragAny reports any fragmentation at all: MF flag or nonzero offset for IPv4, a fragment extension header for IPv6.
|
|
||||||
// Distinct from Packet.Fragment, which is true only for NON-FIRST fragments
|
|
||||||
FragAny bool
|
|
||||||
}
|
|
||||||
|
|||||||
+33
-34
@@ -13,7 +13,6 @@ import (
|
|||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/iputil"
|
|
||||||
"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"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -73,20 +72,20 @@ func TestFirewall_AddRule(t *testing.T) {
|
|||||||
ti6, err := netip.ParsePrefix("fd12::34/128")
|
ti6, err := netip.ParsePrefix("fd12::34/128")
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolTCP, 1, 1, []string{}, "", "", "", "", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoTCP, 1, 1, []string{}, "", "", "", "", ""))
|
||||||
// An empty rule is any
|
// An empty rule is any
|
||||||
assert.True(t, fw.InRules.TCP[1].Any.Any.Any)
|
assert.True(t, fw.InRules.TCP[1].Any.Any.Any)
|
||||||
assert.Empty(t, fw.InRules.TCP[1].Any.Groups)
|
assert.Empty(t, fw.InRules.TCP[1].Any.Groups)
|
||||||
assert.Empty(t, fw.InRules.TCP[1].Any.Hosts)
|
assert.Empty(t, fw.InRules.TCP[1].Any.Hosts)
|
||||||
|
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", ""))
|
||||||
assert.Nil(t, fw.InRules.UDP[1].Any.Any)
|
assert.Nil(t, fw.InRules.UDP[1].Any.Any)
|
||||||
assert.Contains(t, fw.InRules.UDP[1].Any.Groups[0].Groups, "g1")
|
assert.Contains(t, fw.InRules.UDP[1].Any.Groups[0].Groups, "g1")
|
||||||
assert.Empty(t, fw.InRules.UDP[1].Any.Hosts)
|
assert.Empty(t, fw.InRules.UDP[1].Any.Hosts)
|
||||||
|
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 1, 1, []string{}, "h1", "", "", "", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoICMP, 1, 1, []string{}, "h1", "", "", "", ""))
|
||||||
//no matter what port is given for icmp, it should end up as "any"
|
//no matter what port is given for icmp, it should end up as "any"
|
||||||
assert.Nil(t, fw.InRules.ICMP[firewall.PortAny].Any.Any)
|
assert.Nil(t, fw.InRules.ICMP[firewall.PortAny].Any.Any)
|
||||||
assert.Empty(t, fw.InRules.ICMP[firewall.PortAny].Any.Groups)
|
assert.Empty(t, fw.InRules.ICMP[firewall.PortAny].Any.Groups)
|
||||||
@@ -117,11 +116,11 @@ func TestFirewall_AddRule(t *testing.T) {
|
|||||||
assert.True(t, ok)
|
assert.True(t, ok)
|
||||||
|
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", ""))
|
||||||
assert.Contains(t, fw.InRules.UDP[1].CANames, "ca-name")
|
assert.Contains(t, fw.InRules.UDP[1].CANames, "ca-name")
|
||||||
|
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha"))
|
||||||
assert.Contains(t, fw.InRules.UDP[1].CAShas, "ca-sha")
|
assert.Contains(t, fw.InRules.UDP[1].CAShas, "ca-sha")
|
||||||
|
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||||
@@ -186,7 +185,7 @@ func TestFirewall_Drop(t *testing.T) {
|
|||||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||||
LocalPort: 10,
|
LocalPort: 10,
|
||||||
RemotePort: 90,
|
RemotePort: 90,
|
||||||
Protocol: iputil.IPProtocolUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
Fragment: false,
|
Fragment: false,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -264,7 +263,7 @@ func TestFirewall_DropV6(t *testing.T) {
|
|||||||
RemoteAddr: netip.MustParseAddr("fd12::34"),
|
RemoteAddr: netip.MustParseAddr("fd12::34"),
|
||||||
LocalPort: 10,
|
LocalPort: 10,
|
||||||
RemotePort: 90,
|
RemotePort: 90,
|
||||||
Protocol: iputil.IPProtocolUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
Fragment: false,
|
Fragment: false,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -351,7 +350,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
Certificate: &dummyCert{},
|
Certificate: &dummyCert{},
|
||||||
}
|
}
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolUDP}, true, c, cp))
|
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoUDP}, true, c, cp))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -361,7 +360,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
Certificate: &dummyCert{},
|
Certificate: &dummyCert{},
|
||||||
}
|
}
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 1}, true, c, cp))
|
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 1}, true, c, cp))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -371,7 +370,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
}
|
}
|
||||||
ip := netip.MustParsePrefix("9.254.254.254/32")
|
ip := netip.MustParsePrefix("9.254.254.254/32")
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
b.Run("pass proto, port, fail on local CIDRv6", func(b *testing.B) {
|
b.Run("pass proto, port, fail on local CIDRv6", func(b *testing.B) {
|
||||||
@@ -380,7 +379,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
}
|
}
|
||||||
ip := netip.MustParsePrefix("fd99::99/128")
|
ip := netip.MustParsePrefix("fd99::99/128")
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -393,7 +392,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||||
}
|
}
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
|
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
b.Run("pass proto, port, any local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
|
b.Run("pass proto, port, any local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
|
||||||
@@ -405,7 +404,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||||
}
|
}
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
|
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -418,7 +417,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||||
}
|
}
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
|
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
b.Run("pass proto, port, specific local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
|
b.Run("pass proto, port, specific local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
|
||||||
@@ -430,7 +429,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||||
}
|
}
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -442,7 +441,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
InvertedGroups: map[string]struct{}{"good-group": {}},
|
InvertedGroups: map[string]struct{}{"good-group": {}},
|
||||||
}
|
}
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
|
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -454,7 +453,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
InvertedGroups: map[string]struct{}{"good-group": {}},
|
InvertedGroups: map[string]struct{}{"good-group": {}},
|
||||||
}
|
}
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
|
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
b.Run("pass on group on specific local cidr6", func(b *testing.B) {
|
b.Run("pass on group on specific local cidr6", func(b *testing.B) {
|
||||||
@@ -465,7 +464,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
InvertedGroups: map[string]struct{}{"good-group": {}},
|
InvertedGroups: map[string]struct{}{"good-group": {}},
|
||||||
}
|
}
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -477,7 +476,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
|||||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||||
}
|
}
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp)
|
ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -493,7 +492,7 @@ func TestFirewall_Drop2(t *testing.T) {
|
|||||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||||
LocalPort: 10,
|
LocalPort: 10,
|
||||||
RemotePort: 90,
|
RemotePort: 90,
|
||||||
Protocol: iputil.IPProtocolUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
Fragment: false,
|
Fragment: false,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -551,7 +550,7 @@ func TestFirewall_Drop3(t *testing.T) {
|
|||||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||||
LocalPort: 1,
|
LocalPort: 1,
|
||||||
RemotePort: 1,
|
RemotePort: 1,
|
||||||
Protocol: iputil.IPProtocolUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
Fragment: false,
|
Fragment: false,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -639,7 +638,7 @@ func TestFirewall_Drop3V6(t *testing.T) {
|
|||||||
RemoteAddr: netip.MustParseAddr("fd12::34"),
|
RemoteAddr: netip.MustParseAddr("fd12::34"),
|
||||||
LocalPort: 1,
|
LocalPort: 1,
|
||||||
RemotePort: 1,
|
RemotePort: 1,
|
||||||
Protocol: iputil.IPProtocolUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
Fragment: false,
|
Fragment: false,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -676,7 +675,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
|||||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||||
LocalPort: 10,
|
LocalPort: 10,
|
||||||
RemotePort: 90,
|
RemotePort: 90,
|
||||||
Protocol: iputil.IPProtocolUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
Fragment: false,
|
Fragment: false,
|
||||||
}
|
}
|
||||||
network := netip.MustParsePrefix("1.2.3.4/24")
|
network := netip.MustParsePrefix("1.2.3.4/24")
|
||||||
@@ -759,13 +758,13 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
templ := firewall.Packet{
|
templ := firewall.Packet{
|
||||||
LocalAddr: netip.MustParseAddr("1.2.3.4"),
|
LocalAddr: netip.MustParseAddr("1.2.3.4"),
|
||||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||||
Protocol: iputil.IPProtocolICMP,
|
Protocol: firewall.ProtoICMP,
|
||||||
Fragment: false,
|
Fragment: false,
|
||||||
}
|
}
|
||||||
|
|
||||||
t.Run("ICMP allowed", func(t *testing.T) {
|
t.Run("ICMP allowed", func(t *testing.T) {
|
||||||
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 0, 0, []string{"any"}, "", "", "", "", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoICMP, 0, 0, []string{"any"}, "", "", "", "", ""))
|
||||||
t.Run("zero ports", func(t *testing.T) {
|
t.Run("zero ports", func(t *testing.T) {
|
||||||
p := templ.Copy()
|
p := templ.Copy()
|
||||||
p.LocalPort = 0
|
p.LocalPort = 0
|
||||||
@@ -911,7 +910,7 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
|
|||||||
RemoteAddr: netip.MustParseAddr("192.0.2.3"),
|
RemoteAddr: netip.MustParseAddr("192.0.2.3"),
|
||||||
LocalPort: 1,
|
LocalPort: 1,
|
||||||
RemotePort: 1,
|
RemotePort: 1,
|
||||||
Protocol: iputil.IPProtocolUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
Fragment: false,
|
Fragment: false,
|
||||||
}
|
}
|
||||||
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
|
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
|
||||||
@@ -962,7 +961,7 @@ func TestFirewall_ConntrackSourceSpoofingAcrossPeers(t *testing.T) {
|
|||||||
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
||||||
LocalPort: 443,
|
LocalPort: 443,
|
||||||
RemotePort: 55000,
|
RemotePort: 55000,
|
||||||
Protocol: iputil.IPProtocolUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
}
|
}
|
||||||
|
|
||||||
require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil),
|
require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil),
|
||||||
@@ -1032,7 +1031,7 @@ func BenchmarkFirewallDropConntrackHit(b *testing.B) {
|
|||||||
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
||||||
LocalPort: 443,
|
LocalPort: 443,
|
||||||
RemotePort: 55000,
|
RemotePort: 55000,
|
||||||
Protocol: iputil.IPProtocolUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
}
|
}
|
||||||
|
|
||||||
cases := []struct {
|
cases := []struct {
|
||||||
@@ -1318,28 +1317,28 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
|
|||||||
mf := &mockFirewall{}
|
mf := &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "tcp", "host": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "tcp", "host": "a"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding udp rule
|
// Test adding udp rule
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(test.NewLogger())
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "udp", "host": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "udp", "host": "a"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding icmp rule
|
// Test adding icmp rule
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(test.NewLogger())
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "icmp", "host": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "icmp", "host": "a"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding icmp rule no port
|
// Test adding icmp rule no port
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(test.NewLogger())
|
||||||
mf = &mockFirewall{}
|
mf = &mockFirewall{}
|
||||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"proto": "icmp", "host": "a"}}}
|
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"proto": "icmp", "host": "a"}}}
|
||||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||||
|
|
||||||
// Test adding any rule
|
// Test adding any rule
|
||||||
conf = config.NewC(test.NewLogger())
|
conf = config.NewC(test.NewLogger())
|
||||||
@@ -1583,7 +1582,7 @@ func buildTestCase(setup testsetup, err error, theirPrefixes ...netip.Prefix) te
|
|||||||
RemoteAddr: theirPrefixes[0].Addr(),
|
RemoteAddr: theirPrefixes[0].Addr(),
|
||||||
LocalPort: 10,
|
LocalPort: 10,
|
||||||
RemotePort: 90,
|
RemotePort: 90,
|
||||||
Protocol: iputil.IPProtocolUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
Fragment: false,
|
Fragment: false,
|
||||||
}
|
}
|
||||||
return testcase{
|
return testcase{
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ require (
|
|||||||
filippo.io/bigmod v0.1.0
|
filippo.io/bigmod v0.1.0
|
||||||
github.com/anmitsu/go-shlex v0.0.0-20200514113438-38f4b401e2be
|
github.com/anmitsu/go-shlex v0.0.0-20200514113438-38f4b401e2be
|
||||||
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/flynn/noise v1.1.0
|
github.com/flynn/noise v1.1.0
|
||||||
github.com/gaissmai/bart v0.28.0
|
github.com/gaissmai/bart v0.28.0
|
||||||
github.com/gogo/protobuf v1.3.2
|
github.com/gogo/protobuf v1.3.2
|
||||||
@@ -19,7 +20,7 @@ require (
|
|||||||
github.com/rcrowley/go-metrics v0.0.0-20201227073835-cf1acfcdf475
|
github.com/rcrowley/go-metrics v0.0.0-20201227073835-cf1acfcdf475
|
||||||
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e
|
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e
|
||||||
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
||||||
github.com/stretchr/testify v1.12.0
|
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.uber.org/goleak v1.3.0
|
||||||
go.yaml.in/yaml/v3 v3.0.5
|
go.yaml.in/yaml/v3 v3.0.5
|
||||||
@@ -39,8 +40,10 @@ require (
|
|||||||
require (
|
require (
|
||||||
github.com/beorn7/perks v1.0.1 // indirect
|
github.com/beorn7/perks v1.0.1 // indirect
|
||||||
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
||||||
|
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||||
github.com/google/btree v1.1.2 // indirect
|
github.com/google/btree v1.1.2 // indirect
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
|
||||||
|
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||||
github.com/prometheus/client_model v0.6.2 // indirect
|
github.com/prometheus/client_model v0.6.2 // indirect
|
||||||
github.com/prometheus/common v0.70.1 // indirect
|
github.com/prometheus/common v0.70.1 // indirect
|
||||||
github.com/prometheus/procfs v0.21.1 // indirect
|
github.com/prometheus/procfs v0.21.1 // indirect
|
||||||
|
|||||||
@@ -19,7 +19,10 @@ github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6r
|
|||||||
github.com/cespare/xxhash/v2 v2.1.1/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
github.com/cespare/xxhash/v2 v2.1.1/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||||
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
||||||
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||||
|
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432 h1:M5QgkYacWj0Xs8MhpIK/5uwU02icXpEoSo9sM2aRCps=
|
||||||
|
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432/go.mod h1:xwIwAxMvYnVrGJPe2FKx5prTrnAjGOD8zvDOnxnrrkM=
|
||||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
|
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=
|
||||||
@@ -98,6 +101,7 @@ github.com/nbrownus/go-metrics-prometheus v0.0.0-20210712211119-974a6260965f/go.
|
|||||||
github.com/pkg/errors v0.8.0/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
github.com/pkg/errors v0.8.0/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||||
github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||||
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||||
|
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||||
github.com/prometheus/client_golang v0.9.1/go.mod h1:7SWBe2y4D6OKWSNQJUaRYU/AaXPKyh/dDVn+NZz0KFw=
|
github.com/prometheus/client_golang v0.9.1/go.mod h1:7SWBe2y4D6OKWSNQJUaRYU/AaXPKyh/dDVn+NZz0KFw=
|
||||||
github.com/prometheus/client_golang v1.0.0/go.mod h1:db9x61etRT2tGnBNRi70OPL5FsnadC4Ky3P0J6CfImo=
|
github.com/prometheus/client_golang v1.0.0/go.mod h1:db9x61etRT2tGnBNRi70OPL5FsnadC4Ky3P0J6CfImo=
|
||||||
@@ -136,8 +140,8 @@ github.com/stretchr/testify v1.2.2/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXf
|
|||||||
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
||||||
github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4=
|
github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4=
|
||||||
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
||||||
github.com/stretchr/testify v1.12.0 h1:K6Mr6jO9JICuend/5xzTM03ydSV3vdNRYAdPSukj8uI=
|
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
||||||
github.com/stretchr/testify v1.12.0/go.mod h1:bOYBZb5qJ00vPzWfIqBUZPaxK8jWiXc6d3ErP4Ca9Gw=
|
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||||
github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0=
|
github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0=
|
||||||
github.com/vishvananda/netlink v1.3.1/go.mod h1:ARtKouGSTGchR8aMwmkzC0qiNPrrWO5JS/XMVl45+b4=
|
github.com/vishvananda/netlink v1.3.1/go.mod h1:ARtKouGSTGchR8aMwmkzC0qiNPrrWO5JS/XMVl45+b4=
|
||||||
github.com/vishvananda/netns v0.0.5 h1:DfiHV+j8bA32MFM7bfEunvT8IAqQ/NzSJHtcmW5zdEY=
|
github.com/vishvananda/netns v0.0.5 h1:DfiHV+j8bA32MFM7bfEunvT8IAqQ/NzSJHtcmW5zdEY=
|
||||||
|
|||||||
-117
@@ -1,117 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
// This file is a trimmed, inlined copy of the graphite exporter from
|
|
||||||
// github.com/cyberdelia/go-metrics-graphite, retaining only the Config type and
|
|
||||||
// the Once entrypoint that Nebula uses. The upstream package has been
|
|
||||||
// unmaintained for 10+ years, so it was vendored here to drop the dependency.
|
|
||||||
// See https://github.com/slackhq/nebula/issues/1831.
|
|
||||||
//
|
|
||||||
// Copyright 2015 Timothée Peignier. All rights reserved.
|
|
||||||
//
|
|
||||||
// Redistribution and use in source and binary forms, with or without
|
|
||||||
// modification, are permitted provided that the following conditions are met:
|
|
||||||
//
|
|
||||||
// 1. Redistributions of source code must retain the above copyright notice,
|
|
||||||
// this list of conditions and the following disclaimer.
|
|
||||||
//
|
|
||||||
// 2. Redistributions in binary form must reproduce the above copyright notice,
|
|
||||||
// this list of conditions and the following disclaimer in the documentation
|
|
||||||
// and/or other materials provided with the distribution.
|
|
||||||
//
|
|
||||||
// THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
|
||||||
// AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
|
||||||
// IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
|
||||||
// DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
|
||||||
// FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
|
||||||
// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
|
||||||
// SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
|
||||||
// CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
|
||||||
// OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
|
||||||
// OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bufio"
|
|
||||||
"fmt"
|
|
||||||
"net"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
|
||||||
)
|
|
||||||
|
|
||||||
// graphiteConfigExport provides a container with configuration parameters for
|
|
||||||
// the Graphite exporter.
|
|
||||||
type graphiteConfigExport struct {
|
|
||||||
Addr *net.TCPAddr // Network address to connect to
|
|
||||||
Registry metrics.Registry // Registry to be exported
|
|
||||||
FlushInterval time.Duration // Flush interval
|
|
||||||
DurationUnit time.Duration // Time conversion unit for durations
|
|
||||||
Prefix string // Prefix to be prepended to metric names
|
|
||||||
Percentiles []float64 // Percentiles to export from timers and histograms
|
|
||||||
}
|
|
||||||
|
|
||||||
// graphiteOnce performs a single submission to Graphite, returning a non-nil
|
|
||||||
// error on failed connections.
|
|
||||||
func graphiteOnce(c graphiteConfigExport) error {
|
|
||||||
now := time.Now().Unix()
|
|
||||||
du := float64(c.DurationUnit)
|
|
||||||
flushSeconds := float64(c.FlushInterval) / float64(time.Second)
|
|
||||||
conn, err := net.DialTCP("tcp", nil, c.Addr)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
defer conn.Close()
|
|
||||||
w := bufio.NewWriter(conn)
|
|
||||||
c.Registry.Each(func(name string, i any) {
|
|
||||||
switch metric := i.(type) {
|
|
||||||
case metrics.Counter:
|
|
||||||
count := metric.Count()
|
|
||||||
fmt.Fprintf(w, "%s.%s.count %d %d\n", c.Prefix, name, count, now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.count_ps %.2f %d\n", c.Prefix, name, float64(count)/flushSeconds, now)
|
|
||||||
case metrics.Gauge:
|
|
||||||
fmt.Fprintf(w, "%s.%s.value %d %d\n", c.Prefix, name, metric.Value(), now)
|
|
||||||
case metrics.GaugeFloat64:
|
|
||||||
fmt.Fprintf(w, "%s.%s.value %f %d\n", c.Prefix, name, metric.Value(), now)
|
|
||||||
case metrics.Histogram:
|
|
||||||
h := metric.Snapshot()
|
|
||||||
ps := h.Percentiles(c.Percentiles)
|
|
||||||
fmt.Fprintf(w, "%s.%s.count %d %d\n", c.Prefix, name, h.Count(), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.min %d %d\n", c.Prefix, name, h.Min(), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.max %d %d\n", c.Prefix, name, h.Max(), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.mean %.2f %d\n", c.Prefix, name, h.Mean(), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.std-dev %.2f %d\n", c.Prefix, name, h.StdDev(), now)
|
|
||||||
for psIdx, psKey := range c.Percentiles {
|
|
||||||
key := strings.Replace(strconv.FormatFloat(psKey*100.0, 'f', -1, 64), ".", "", 1)
|
|
||||||
fmt.Fprintf(w, "%s.%s.%s-percentile %.2f %d\n", c.Prefix, name, key, ps[psIdx], now)
|
|
||||||
}
|
|
||||||
case metrics.Meter:
|
|
||||||
m := metric.Snapshot()
|
|
||||||
fmt.Fprintf(w, "%s.%s.count %d %d\n", c.Prefix, name, m.Count(), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.one-minute %.2f %d\n", c.Prefix, name, m.Rate1(), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.five-minute %.2f %d\n", c.Prefix, name, m.Rate5(), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.fifteen-minute %.2f %d\n", c.Prefix, name, m.Rate15(), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.mean %.2f %d\n", c.Prefix, name, m.RateMean(), now)
|
|
||||||
case metrics.Timer:
|
|
||||||
t := metric.Snapshot()
|
|
||||||
ps := t.Percentiles(c.Percentiles)
|
|
||||||
count := t.Count()
|
|
||||||
fmt.Fprintf(w, "%s.%s.count %d %d\n", c.Prefix, name, count, now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.count_ps %.2f %d\n", c.Prefix, name, float64(count)/flushSeconds, now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.min %d %d\n", c.Prefix, name, t.Min()/int64(du), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.max %d %d\n", c.Prefix, name, t.Max()/int64(du), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.mean %.2f %d\n", c.Prefix, name, t.Mean()/du, now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.std-dev %.2f %d\n", c.Prefix, name, t.StdDev()/du, now)
|
|
||||||
for psIdx, psKey := range c.Percentiles {
|
|
||||||
key := strings.Replace(strconv.FormatFloat(psKey*100.0, 'f', -1, 64), ".", "", 1)
|
|
||||||
fmt.Fprintf(w, "%s.%s.%s-percentile %.2f %d\n", c.Prefix, name, key, ps[psIdx]/du, now)
|
|
||||||
}
|
|
||||||
fmt.Fprintf(w, "%s.%s.one-minute %.2f %d\n", c.Prefix, name, t.Rate1(), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.five-minute %.2f %d\n", c.Prefix, name, t.Rate5(), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.fifteen-minute %.2f %d\n", c.Prefix, name, t.Rate15(), now)
|
|
||||||
fmt.Fprintf(w, "%s.%s.mean-rate %.2f %d\n", c.Prefix, name, t.RateMean(), now)
|
|
||||||
}
|
|
||||||
w.Flush()
|
|
||||||
})
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
+3
-18
@@ -749,14 +749,8 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
connState, err := newConnectionStateFromResult(result)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Discarding handshake with an invalid message index", "error", err, "vpnAddrs", vpnAddrs)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo := &HostInfo{
|
hostinfo := &HostInfo{
|
||||||
ConnectionState: connState,
|
ConnectionState: newConnectionStateFromResult(result),
|
||||||
localIndexId: result.LocalIndex,
|
localIndexId: result.LocalIndex,
|
||||||
remoteIndexId: result.RemoteIndex,
|
remoteIndexId: result.RemoteIndex,
|
||||||
vpnAddrs: vpnAddrs,
|
vpnAddrs: vpnAddrs,
|
||||||
@@ -874,13 +868,7 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Handshake complete; build the ConnectionState now that we have keys and a verified peer cert.
|
// Handshake complete; build the ConnectionState now that we have keys and a verified peer cert.
|
||||||
cs, err := newConnectionStateFromResult(result)
|
hostinfo.ConnectionState = newConnectionStateFromResult(result)
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Discarding handshake with an invalid message index", "error", err, "vpnAddrs", hostinfo.vpnAddrs)
|
|
||||||
hm.DeleteHostInfo(hostinfo)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
hostinfo.ConnectionState = cs
|
|
||||||
|
|
||||||
remoteCert := result.RemoteCert
|
remoteCert := result.RemoteCert
|
||||||
if remoteCert == nil {
|
if remoteCert == nil {
|
||||||
@@ -987,9 +975,6 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
|||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
out := make([]byte, mtu)
|
out := make([]byte, mtu)
|
||||||
for _, cp := range hh.packetStore {
|
for _, cp := range hh.packetStore {
|
||||||
// TODO: use a SendBatch here. Each callback lands in
|
|
||||||
// sendNoMetrics -> WriteTo: one syscall per cached packet,
|
|
||||||
// where one sendmmsg could flush the whole store.
|
|
||||||
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
||||||
}
|
}
|
||||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||||
@@ -1101,7 +1086,7 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos
|
|||||||
// We received a valid handshake on this relay, so make sure the relay
|
// We received a valid handshake on this relay, so make sure the relay
|
||||||
// state reflects that, in case it had been marked Disestablished.
|
// state reflects that, in case it had been marked Disestablished.
|
||||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established)
|
via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established)
|
||||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false, 0)
|
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||||
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
|
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -84,7 +84,7 @@ func (mw *mockEncWriter) SendMessageToVpnAddr(_ header.MessageType, _ header.Mes
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
func (mw *mockEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
func (mw *mockEncWriter) SendVia(_ *HostInfo, _ *Relay, _, _, _ []byte, _ bool) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+6
-11
@@ -190,18 +190,13 @@ func SubTypeName(t MessageType, s MessageSubType) string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func IsValidSubType(t MessageType, s MessageSubType) bool {
|
func IsValidSubType(t MessageType, s MessageSubType) bool {
|
||||||
switch t {
|
if n, ok := subTypeMap[t]; ok {
|
||||||
case Message:
|
if _, ok := (*n)[s]; ok {
|
||||||
return s == MessageNone || s == MessageRelay
|
return true
|
||||||
case Handshake:
|
}
|
||||||
return s == HandshakeIXPSK0
|
|
||||||
case Test:
|
|
||||||
return s == TestReply || s == TestRequest
|
|
||||||
case Control, CloseTunnel, RecvError, LightHouse:
|
|
||||||
return s == 0
|
|
||||||
default:
|
|
||||||
return false
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewHeader turns bytes into a header
|
// NewHeader turns bytes into a header
|
||||||
|
|||||||
@@ -102,57 +102,6 @@ func TestTypeMap(t *testing.T) {
|
|||||||
}, subTypeMap)
|
}, subTypeMap)
|
||||||
}
|
}
|
||||||
|
|
||||||
// mapIsValidSubType is the pre-refactor, map-driven definition of a valid
|
|
||||||
// subtype. IsValidSubType was reimplemented as an explicit switch; this keeps
|
|
||||||
// the original behavior around so we can prove the switch is equivalent to it.
|
|
||||||
func mapIsValidSubType(t MessageType, s MessageSubType) bool {
|
|
||||||
if n, ok := subTypeMap[t]; ok {
|
|
||||||
if _, ok := (*n)[s]; ok {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestIsValidSubType(t *testing.T) {
|
|
||||||
// Explicit intent table: documents exactly which subtypes are valid so the
|
|
||||||
// test stays meaningful even if both the switch and subTypeMap change.
|
|
||||||
assert.True(t, IsValidSubType(Message, MessageNone))
|
|
||||||
assert.True(t, IsValidSubType(Message, MessageRelay))
|
|
||||||
assert.False(t, IsValidSubType(Message, 2))
|
|
||||||
|
|
||||||
assert.True(t, IsValidSubType(Handshake, HandshakeIXPSK0))
|
|
||||||
// HandshakeXXPSK0 is defined but not a wire-valid subtype.
|
|
||||||
assert.False(t, IsValidSubType(Handshake, HandshakeXXPSK0))
|
|
||||||
|
|
||||||
assert.True(t, IsValidSubType(Test, TestRequest))
|
|
||||||
assert.True(t, IsValidSubType(Test, TestReply))
|
|
||||||
assert.False(t, IsValidSubType(Test, 2))
|
|
||||||
|
|
||||||
// These types only ever carry subtype 0.
|
|
||||||
for _, mt := range []MessageType{Control, CloseTunnel, RecvError, LightHouse} {
|
|
||||||
assert.True(t, IsValidSubType(mt, 0), "type %d subtype 0 should be valid", mt)
|
|
||||||
assert.False(t, IsValidSubType(mt, 1), "type %d subtype 1 should be invalid", mt)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Unknown/unassigned types are never valid.
|
|
||||||
assert.False(t, IsValidSubType(99, 0))
|
|
||||||
|
|
||||||
// Exhaustive proof of equivalence with the original map-driven logic across
|
|
||||||
// the entire (type, subtype) input space.
|
|
||||||
for ti := 0; ti <= 0xff; ti++ {
|
|
||||||
for si := 0; si <= 0xff; si++ {
|
|
||||||
mt, mst := MessageType(ti), MessageSubType(si)
|
|
||||||
assert.Equalf(t, mapIsValidSubType(mt, mst), IsValidSubType(mt, mst),
|
|
||||||
"IsValidSubType(%d, %d) diverged from map-driven definition", ti, si)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// H method must delegate to the package function.
|
|
||||||
assert.True(t, (&H{Type: Test, Subtype: TestReply}).IsValidSubType())
|
|
||||||
assert.False(t, (&H{Type: Handshake, Subtype: HandshakeXXPSK0}).IsValidSubType())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHeader_String(t *testing.T) {
|
func TestHeader_String(t *testing.T) {
|
||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
|
|||||||
+13
-78
@@ -239,15 +239,11 @@ const (
|
|||||||
|
|
||||||
type HostInfo struct {
|
type HostInfo struct {
|
||||||
remote atomic.Pointer[netip.AddrPort]
|
remote atomic.Pointer[netip.AddrPort]
|
||||||
|
remotes *RemoteList
|
||||||
|
promoteCounter atomic.Uint32
|
||||||
ConnectionState *ConnectionState
|
ConnectionState *ConnectionState
|
||||||
|
remoteIndexId uint32
|
||||||
// Traffic bits, pendingDeletion, and the rebind epoch we last sent under
|
localIndexId uint32
|
||||||
state atomic.Uint32
|
|
||||||
|
|
||||||
promoteCounter atomic.Uint32
|
|
||||||
remoteIndexId uint32
|
|
||||||
localIndexId uint32
|
|
||||||
remotes *RemoteList
|
|
||||||
|
|
||||||
// vpnAddrs is a list of vpn addresses assigned to this host that are within our own vpn networks
|
// vpnAddrs is a list of vpn addresses assigned to this host that are within our own vpn networks
|
||||||
// The host may have other vpn addresses that are outside our
|
// The host may have other vpn addresses that are outside our
|
||||||
@@ -266,6 +262,11 @@ type HostInfo struct {
|
|||||||
// This is used to limit lighthouse re-queries in chatty clients
|
// This is used to limit lighthouse re-queries in chatty clients
|
||||||
nextLHQuery atomic.Int64
|
nextLHQuery atomic.Int64
|
||||||
|
|
||||||
|
// lastRebindCount is the other side of Interface.rebindCount, if these values don't match then we need to ask LH
|
||||||
|
// for a punch from the remote end of this tunnel. The goal being to prime their conntrack for our traffic just like
|
||||||
|
// with a handshake
|
||||||
|
lastRebindCount int8
|
||||||
|
|
||||||
// lastHandshakeTime records the time the remote side told us about at the stage when the handshake was completed locally
|
// lastHandshakeTime records the time the remote side told us about at the stage when the handshake was completed locally
|
||||||
// Stage 1 packet will contain it if I am a responder, stage 2 packet if I am an initiator
|
// Stage 1 packet will contain it if I am a responder, stage 2 packet if I am an initiator
|
||||||
// This is used to avoid an attack where a handshake packet is replayed after some time
|
// This is used to avoid an attack where a handshake packet is replayed after some time
|
||||||
@@ -274,6 +275,9 @@ type HostInfo struct {
|
|||||||
lastRoam time.Time
|
lastRoam time.Time
|
||||||
lastRoamRemote netip.AddrPort
|
lastRoamRemote netip.AddrPort
|
||||||
|
|
||||||
|
//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
|
||||||
|
|
||||||
// lastUsed tracks the last time ConnectionManager checked the tunnel and it was in use.
|
// lastUsed tracks the last time ConnectionManager checked the tunnel and it was in use.
|
||||||
// This value will be behind against actual tunnel utilization in the hot path.
|
// This value will be behind against actual tunnel utilization in the hot path.
|
||||||
// This should only be used by the ConnectionManagers ticker routine.
|
// This should only be used by the ConnectionManagers ticker routine.
|
||||||
@@ -539,17 +543,6 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
|||||||
return final
|
return final
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HostMap) QueryIndexCached(index uint32, cache map[uint32]*HostInfo) *HostInfo {
|
|
||||||
if out, ok := cache[index]; ok {
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
out := hm.QueryIndex(index)
|
|
||||||
if out != nil {
|
|
||||||
cache[index] = out
|
|
||||||
}
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
|
func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
|
||||||
hm.RLock()
|
hm.RLock()
|
||||||
if h, ok := hm.Indexes[index]; ok {
|
if h, ok := hm.Indexes[index]; ok {
|
||||||
@@ -665,7 +658,7 @@ 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.markOut(f.rebindEpoch.Load())
|
hostinfo.out.Store(true)
|
||||||
if f.connectionManager != nil { // f.connectionManager is only nil in some unit tests
|
if f.connectionManager != nil { // f.connectionManager is only nil in some unit tests
|
||||||
f.connectionManager.trafficTimer.Add(hostinfo.localIndexId, f.connectionManager.checkInterval)
|
f.connectionManager.trafficTimer.Add(hostinfo.localIndexId, f.connectionManager.checkInterval)
|
||||||
}
|
}
|
||||||
@@ -766,64 +759,6 @@ func (i *HostInfo) TryPromoteBest(preferredRanges []netip.Prefix, ifce *Interfac
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Bits within HostInfo.state, everything above stateEpochShift is the epoch
|
|
||||||
const (
|
|
||||||
stateIn uint32 = 1 << iota
|
|
||||||
stateOut
|
|
||||||
statePendingDeletion
|
|
||||||
|
|
||||||
stateFlags = stateIn | stateOut | statePendingDeletion
|
|
||||||
// The epoch is the top 29 bits, it would take 2^29 rebinds to wrap and we will never get there
|
|
||||||
stateEpochShift = 3
|
|
||||||
)
|
|
||||||
|
|
||||||
// markIn records inbound traffic
|
|
||||||
func (i *HostInfo) markIn() {
|
|
||||||
if i.state.Load()&stateIn == 0 {
|
|
||||||
i.state.Or(stateIn)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// markOut records a send and reports whether the epoch moved, meaning we want a punch from the far side
|
|
||||||
func (i *HostInfo) markOut(epoch uint32) bool {
|
|
||||||
e := epoch << stateEpochShift
|
|
||||||
for {
|
|
||||||
old := i.state.Load()
|
|
||||||
if old&stateOut != 0 && old&^stateFlags == e {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
if i.state.CompareAndSwap(old, old&stateFlags|stateOut|e) {
|
|
||||||
return old&^stateFlags != e
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// markOutOnly records a send without consuming the rebind epoch, for paths that cannot act on a requery
|
|
||||||
func (i *HostInfo) markOutOnly() {
|
|
||||||
if i.state.Load()&stateOut == 0 {
|
|
||||||
i.state.Or(stateOut)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// takeTraffic clears both traffic bits, leaving the epoch alone, and reports what they were
|
|
||||||
func (i *HostInfo) takeTraffic() (in bool, out bool) {
|
|
||||||
old := i.state.And(^(stateIn | stateOut))
|
|
||||||
return old&stateIn != 0, old&stateOut != 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (i *HostInfo) setPendingDeletion(v bool) {
|
|
||||||
if v {
|
|
||||||
i.state.Or(statePendingDeletion)
|
|
||||||
} else {
|
|
||||||
i.state.And(^statePendingDeletion)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (i *HostInfo) isPendingDeletion() bool {
|
|
||||||
return i.state.Load()&statePendingDeletion != 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (i *HostInfo) GetCert() *cert.CachedCertificate {
|
func (i *HostInfo) GetCert() *cert.CachedCertificate {
|
||||||
if i.ConnectionState != nil {
|
if i.ConnectionState != nil {
|
||||||
return i.ConnectionState.peerCert
|
return i.ConnectionState.peerCert
|
||||||
|
|||||||
@@ -401,49 +401,3 @@ func TestHostMap_RelayState(t *testing.T) {
|
|||||||
assert.Equal(t, []netip.Addr{}, h1.relayState.relays)
|
assert.Equal(t, []netip.Addr{}, h1.relayState.relays)
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// sentSinceCheck reports whether anything has been sent since the connection manager last looked. Test only:
|
|
||||||
// production reads the out bit through takeTraffic on the connection manager tick.
|
|
||||||
func (i *HostInfo) sentSinceCheck() bool {
|
|
||||||
return i.state.Load()&stateOut != 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHostInfo_markOut(t *testing.T) {
|
|
||||||
h := &HostInfo{}
|
|
||||||
h.markOut(5) // stamped when the tunnel was added
|
|
||||||
|
|
||||||
// A tunnel already on the current epoch has nothing to report, which is what keeps a fresh tunnel from
|
|
||||||
// requerying on its first packet
|
|
||||||
assert.False(t, h.markOut(5), "an unchanged epoch should not report a move")
|
|
||||||
assert.True(t, h.sentSinceCheck(), "the send is still recorded as traffic")
|
|
||||||
|
|
||||||
// A rebind is observed exactly once, so we requery once per rebind
|
|
||||||
assert.True(t, h.markOut(6), "a bumped epoch should report a move")
|
|
||||||
assert.False(t, h.markOut(6), "the epoch move should only be reported once")
|
|
||||||
|
|
||||||
// Traffic and pendingDeletion live in the same word and must survive an epoch change
|
|
||||||
h.setPendingDeletion(true)
|
|
||||||
h.markIn()
|
|
||||||
assert.True(t, h.markOut(7))
|
|
||||||
assert.True(t, h.isPendingDeletion(), "pendingDeletion must survive an epoch change")
|
|
||||||
in, out := h.takeTraffic()
|
|
||||||
assert.True(t, in, "inbound traffic must survive an epoch change")
|
|
||||||
assert.True(t, out)
|
|
||||||
|
|
||||||
// Clearing the traffic bits leaves the epoch alone, otherwise an idle tunnel would requery forever
|
|
||||||
assert.False(t, h.markOut(7), "takeTraffic must not disturb the epoch")
|
|
||||||
}
|
|
||||||
|
|
||||||
// A relayed send records traffic but must leave the rebind epoch for the direct path to consume, otherwise
|
|
||||||
// relaying to a host swallows the requery that gets the far side punching at our new address.
|
|
||||||
func TestHostInfo_markOutOnly(t *testing.T) {
|
|
||||||
h := &HostInfo{}
|
|
||||||
h.markOut(5)
|
|
||||||
|
|
||||||
h.markOutOnly()
|
|
||||||
assert.True(t, h.sentSinceCheck(), "a relayed send is still outbound traffic")
|
|
||||||
assert.False(t, h.markOut(5), "a relayed send must not disturb the epoch")
|
|
||||||
|
|
||||||
assert.True(t, h.markOut(6), "a relayed send must not consume the epoch edge")
|
|
||||||
assert.False(t, h.markOut(6))
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -2,8 +2,6 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
@@ -11,24 +9,10 @@ import (
|
|||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
"github.com/slackhq/nebula/overlay/batch"
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.ParsedPacket, nb []byte, sendBatch *batch.SendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
|
|
||||||
// only valid until the next Read on that queue. Every consumer below
|
|
||||||
// (parse, self-forward, handshake cache, sendInsideMessage) reads it
|
|
||||||
// synchronously; do not retain pkt outside this call. If a future
|
|
||||||
// caller needs to keep the packet, use pkt.Clone() to detach it from
|
|
||||||
// the borrow.
|
|
||||||
//
|
|
||||||
// pkt.Bytes is either one IP datagram (GSO zero) or a TSO/USO
|
|
||||||
// superpacket. In both cases the L3+L4 headers at the start describe
|
|
||||||
// the same 5-tuple every segment will share, so a single newPacket /
|
|
||||||
// firewall check covers the whole superpacket.
|
|
||||||
packet := pkt.Bytes
|
|
||||||
err := newPacket(packet, false, fwPacket)
|
err := newPacket(packet, false, fwPacket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
@@ -53,17 +37,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
|
|||||||
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
||||||
// TUN device.
|
// TUN device.
|
||||||
if immediatelyForwardToSelf {
|
if immediatelyForwardToSelf {
|
||||||
// Write copies into the kernel queue synchronously, so seg's lifetime ends at return.
|
_, err := f.readers[q].Write(packet)
|
||||||
// A self-forwarded superpacket would be re-handed to the
|
|
||||||
// kernel as one giant blob; segment first so the loopback
|
|
||||||
// path sees one IP datagram per Write.
|
|
||||||
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
|
||||||
// The kernel may have left the transport checksum for hardware
|
|
||||||
// offload to finish; nothing between here and the tun will.
|
|
||||||
iputil.SetTransportChecksum(seg)
|
|
||||||
_, werr := f.queues[q].Write(seg)
|
|
||||||
return werr
|
|
||||||
})
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to forward to tun", "error", err)
|
f.l.Error("Failed to forward to tun", "error", err)
|
||||||
}
|
}
|
||||||
@@ -78,24 +52,12 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(&fwPacket.Packet, func(hh *HandshakeHostInfo) {
|
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
||||||
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
|
||||||
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
|
||||||
// so retaining segments past the loop is safe.
|
|
||||||
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
|
||||||
hh.cachePacket(f.l, header.Message, 0, seg, f.sendMessageNow, f.cachedPacketMetrics)
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
if err != nil && f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
f.l.Debug("Failed to segment superpacket for handshake cache",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddr", fwPacket.RemoteAddr,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
})
|
})
|
||||||
|
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
f.rejectInside(packet, rejectBuf, q)
|
f.rejectInside(packet, out, 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,
|
||||||
@@ -109,11 +71,12 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
dropReason := f.firewall.Drop(fwPacket.Packet, false, hostinfo, f.pki.GetCAPool(), localCache)
|
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
if dropReason == nil {
|
if dropReason == nil {
|
||||||
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch)
|
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, rejectBuf, q)
|
f.rejectInside(packet, out, 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,
|
||||||
@@ -123,122 +86,6 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, seg, scratch, nb []byte) []byte {
|
|
||||||
if noiseutil.EncryptLockNeeded {
|
|
||||||
ci.writeLock.Lock()
|
|
||||||
}
|
|
||||||
c := ci.messageCounter.Add(1)
|
|
||||||
|
|
||||||
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
|
||||||
|
|
||||||
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
|
|
||||||
if noiseutil.EncryptLockNeeded {
|
|
||||||
ci.writeLock.Unlock()
|
|
||||||
}
|
|
||||||
if encErr != nil {
|
|
||||||
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
|
|
||||||
"error", encErr,
|
|
||||||
"udpAddr", hostinfo.GetRemote(),
|
|
||||||
"counter", c,
|
|
||||||
)
|
|
||||||
// Skip this segment; the rest of the superpacket can still go out. TCP will retransmit anything we drop here.
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
// sendInsideMessage encrypts a firewall-approved inside packet (or every
|
|
||||||
// segment of a TSO/USO superpacket) into the caller's batch slot for
|
|
||||||
// later sendmmsg flush. Segmentation is fused with encryption here so the
|
|
||||||
// kernel-supplied superpacket bytes never get written into a separate
|
|
||||||
// scratch arena: SegmentSuperpacket builds each segment's plaintext in
|
|
||||||
// segScratch[:segLen] in turn, and we encrypt directly into a fresh SendBatch slot.
|
|
||||||
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch *batch.SendBatch) {
|
|
||||||
ci := hostinfo.ConnectionState
|
|
||||||
if ci.eKey == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// One traffic-out mark covers every segment of the superpacket; doing it
|
|
||||||
// per segment in sendInsideEncrypt paid an atomic store up to ~45 extra
|
|
||||||
// times per TSO packet, inside writeLock under boring crypto.
|
|
||||||
//
|
|
||||||
// We rebound since this tunnel last sent, ask the lighthouse to get the far side punching at us again
|
|
||||||
if f.connectionManager.Out(hostinfo) {
|
|
||||||
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind epoch",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
remote := hostinfo.GetRemote()
|
|
||||||
if !remote.IsValid() { //the relay path
|
|
||||||
//first, find our relay hostinfo:
|
|
||||||
var relayHostInfo *HostInfo
|
|
||||||
var relay *Relay
|
|
||||||
var err error
|
|
||||||
for _, relayIP := range hostinfo.relayState.CopyRelayIps() {
|
|
||||||
relayHostInfo, relay, err = f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP)
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.relayState.DeleteRelay(relayIP)
|
|
||||||
hostinfo.logger(f.l).Info("sendNoMetrics failed to find HostInfo",
|
|
||||||
"relay", relayIP,
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
break
|
|
||||||
}
|
|
||||||
if relayHostInfo == nil || relay == nil {
|
|
||||||
//failure already logged
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
err = tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
|
||||||
//relay header + header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305) + relay tag
|
|
||||||
scratch := sendBatch.Reserve(header.Len + header.Len + len(seg) + 16 + 16)
|
|
||||||
|
|
||||||
innerPacket := f.sendInsideEncrypt(hostinfo, ci, seg, scratch[header.Len:], nb)
|
|
||||||
if innerPacket == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
//now we need to do a relay-encrypt:
|
|
||||||
toSend, err := f.prepareSendVia(relayHostInfo, relay, innerPacket, nb, scratch, true)
|
|
||||||
if err != nil {
|
|
||||||
//already logged
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
sendBatch.Commit(toSend, relayHostInfo.GetRemote())
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(f.l).Error("Failed to segment superpacket for relay send", "error", err)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
|
||||||
// header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305)
|
|
||||||
scratch := sendBatch.Reserve(header.Len + len(seg) + 16)
|
|
||||||
|
|
||||||
out := f.sendInsideEncrypt(hostinfo, ci, seg, scratch, nb)
|
|
||||||
if out == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
sendBatch.Commit(out, remote)
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(f.l).Error("Failed to segment superpacket for send", "error", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||||
if !f.firewall.OutboundSendReject {
|
if !f.firewall.OutboundSendReject {
|
||||||
return
|
return
|
||||||
@@ -249,36 +96,33 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err := f.queues[q].Write(out)
|
_, err := f.readers[q].Write(out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to write to tun", "error", err)
|
f.l.Error("Failed to write to tun", "error", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, rejectBuf []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.InboundSendReject {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// split rejectBuf to make sure we have room to write the plaintext rejection, then encrypt it, without trampling anything
|
out = iputil.CreateRejectPacket(packet, out)
|
||||||
// we can't re-use packet, if we need to send an icmp reject, it won't be long enough.
|
|
||||||
half := len(rejectBuf) / 2
|
|
||||||
encryptBuf := rejectBuf[0:0:half] //the first half of rejectBuf's capacity, len set to 0
|
|
||||||
buildBuf := rejectBuf[half:]
|
|
||||||
|
|
||||||
out := iputil.CreateRejectPacket(packet, buildBuf)
|
|
||||||
if len(out) == 0 {
|
if len(out) == 0 {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(out) > iputil.MaxRejectPacketSize {
|
if len(out) > iputil.MaxRejectPacketSize {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelInfo) {
|
if f.l.Enabled(context.Background(), slog.LevelInfo) {
|
||||||
f.l.Info("rejectOutside: packet too big, not sending", "packet", packet, "outPacket", out)
|
f.l.Info("rejectOutside: packet too big, not sending",
|
||||||
|
"packet", packet,
|
||||||
|
"outPacket", out,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, encryptBuf, q)
|
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, packet, q)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Handshake will attempt to initiate a tunnel with the provided vpn address. This is a no-op if the tunnel is already established or being established
|
// Handshake will attempt to initiate a tunnel with the provided vpn address. This is a no-op if the tunnel is already established or being established
|
||||||
@@ -372,7 +216,7 @@ func (f *Interface) getOrHandshakeConsiderRouting(fwPacket *firewall.Packet, cac
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) {
|
func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) {
|
||||||
fp := &firewall.ParsedPacket{}
|
fp := &firewall.Packet{}
|
||||||
err := newPacket(p, false, fp)
|
err := newPacket(p, false, fp)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
|
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
|
||||||
@@ -380,7 +224,7 @@ func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubTyp
|
|||||||
}
|
}
|
||||||
|
|
||||||
// check if packet is in outbound fw rules
|
// check if packet is in outbound fw rules
|
||||||
dropReason := f.firewall.Drop(fp.Packet, false, hostinfo, f.pki.GetCAPool(), nil)
|
dropReason := f.firewall.Drop(*fp, false, hostinfo, f.pki.GetCAPool(), nil)
|
||||||
if dropReason != nil {
|
if dropReason != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("dropping cached packet",
|
f.l.Debug("dropping cached packet",
|
||||||
@@ -431,36 +275,29 @@ func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *C
|
|||||||
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
// dropExhausted records an exhaustion drop and logs once, on the crossing send, for a spent tunnel.
|
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
||||||
func (f *Interface) dropExhausted(hostinfo *HostInfo, c uint64, msg string) {
|
// to the payload for the ultimate target host, making this a useful method for sending
|
||||||
f.messageMetrics.TxExhausted(1)
|
// handshake messages to peers through relay tunnels.
|
||||||
if c == RejectAfterMessages {
|
// via is the HostInfo through which the message is relayed.
|
||||||
hostinfo.logger(f.l).Error(msg)
|
// ad is the plaintext data to authenticate, but not encrypt
|
||||||
}
|
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
||||||
}
|
// out is a buffer used to store the result of the Encrypt operation
|
||||||
|
// q indicates which writer to use to send the packet.
|
||||||
func (f *Interface) prepareSendVia(via *HostInfo,
|
func (f *Interface) SendVia(via *HostInfo,
|
||||||
relay *Relay,
|
relay *Relay,
|
||||||
ad,
|
ad,
|
||||||
nb,
|
nb,
|
||||||
out []byte,
|
out []byte,
|
||||||
nocopy bool,
|
nocopy bool,
|
||||||
) ([]byte, error) {
|
) {
|
||||||
if noiseutil.EncryptLockNeeded {
|
if noiseutil.EncryptLockNeeded {
|
||||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||||
via.ConnectionState.writeLock.Lock()
|
via.ConnectionState.writeLock.Lock()
|
||||||
}
|
}
|
||||||
c, ok := via.ConnectionState.NextMessageCounter()
|
c := via.ConnectionState.messageCounter.Add(1)
|
||||||
if !ok {
|
|
||||||
if noiseutil.EncryptLockNeeded {
|
|
||||||
via.ConnectionState.writeLock.Unlock()
|
|
||||||
}
|
|
||||||
f.dropExhausted(via, c, "Dropping outbound relay packets, tunnel message counter is exhausted")
|
|
||||||
return nil, fmt.Errorf("tunnel message counter is exhausted")
|
|
||||||
}
|
|
||||||
|
|
||||||
out = header.Encode(out, header.Version, header.Message, header.MessageRelay, relay.RemoteIndex, c)
|
out = header.Encode(out, header.Version, header.Message, header.MessageRelay, relay.RemoteIndex, c)
|
||||||
f.connectionManager.OutNoRebind(via)
|
f.connectionManager.Out(via)
|
||||||
|
|
||||||
// Authenticate the header and payload, but do not encrypt for this message type.
|
// Authenticate the header and payload, but do not encrypt for this message type.
|
||||||
// The payload consists of the inner, unencrypted Nebula header, as well as the end-to-end encrypted payload.
|
// The payload consists of the inner, unencrypted Nebula header, as well as the end-to-end encrypted payload.
|
||||||
@@ -474,7 +311,7 @@ func (f *Interface) prepareSendVia(via *HostInfo,
|
|||||||
"headerLen", len(out),
|
"headerLen", len(out),
|
||||||
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
||||||
)
|
)
|
||||||
return nil, io.ErrShortBuffer
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
||||||
@@ -494,31 +331,13 @@ func (f *Interface) prepareSendVia(via *HostInfo,
|
|||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
f.connectionManager.RelayUsed(relay.LocalIndex)
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
|
||||||
// to the payload for the ultimate target host, making this a useful method for sending
|
|
||||||
// handshake messages to peers through relay tunnels.
|
|
||||||
// via is the HostInfo through which the message is relayed.
|
|
||||||
// ad is the plaintext data to authenticate, but not encrypt
|
|
||||||
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
|
||||||
// out is a buffer used to store the result of the Encrypt operation
|
|
||||||
// q indicates which writer to use to send the packet.
|
|
||||||
func (f *Interface) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
|
||||||
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
|
|
||||||
if err != nil {
|
|
||||||
// already logged by prepareSendVia
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
err = f.writers[0].WriteTo(out, via.GetRemote())
|
||||||
err = f.writers[q].WriteTo(toSend, via.GetRemote())
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
||||||
}
|
}
|
||||||
|
f.connectionManager.RelayUsed(relay.LocalIndex)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
||||||
@@ -542,24 +361,21 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
|||||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||||
ci.writeLock.Lock()
|
ci.writeLock.Lock()
|
||||||
}
|
}
|
||||||
c, ok := ci.NextMessageCounter()
|
c := ci.messageCounter.Add(1)
|
||||||
if !ok {
|
|
||||||
if noiseutil.EncryptLockNeeded {
|
|
||||||
ci.writeLock.Unlock()
|
|
||||||
}
|
|
||||||
f.dropExhausted(hostinfo, c, "Dropping outbound packets, tunnel message counter is exhausted")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
//l.WithField("trace", string(debug.Stack())).Error("out Header ", &Header{Version, t, st, 0, hostinfo.remoteIndexId, c}, p)
|
//l.WithField("trace", string(debug.Stack())).Error("out Header ", &Header{Version, t, st, 0, hostinfo.remoteIndexId, c}, p)
|
||||||
out = header.Encode(out, header.Version, t, st, hostinfo.remoteIndexId, c)
|
out = header.Encode(out, header.Version, t, st, hostinfo.remoteIndexId, c)
|
||||||
// A closing tunnel is torn down right after this, so skip the connection manager entirely: no point recording
|
f.connectionManager.Out(hostinfo)
|
||||||
// traffic or asking the lighthouse for a punch. Otherwise, if we rebound since this tunnel last sent, ask the
|
|
||||||
// lighthouse to get the far side punching at us again.
|
// Query our LH if we haven't since the last time we've been rebound, this will cause the remote to punch against
|
||||||
if t != header.CloseTunnel && f.connectionManager.Out(hostinfo) {
|
// all our addrs and enable a faster roaming.
|
||||||
|
if t != header.CloseTunnel && 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])
|
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
|
||||||
|
hostinfo.lastRebindCount = f.rebindCount
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("Lighthouse update triggered for punch due to rebind epoch",
|
f.l.Debug("Lighthouse update triggered for punch due to rebind counter",
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -607,7 +423,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
|||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true, q)
|
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true)
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
-265
@@ -1,265 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"io"
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
|
||||||
"github.com/slackhq/nebula/iputil"
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
ipv4HeaderLen = 20
|
|
||||||
ipv6HeaderLen = 40
|
|
||||||
)
|
|
||||||
|
|
||||||
// capturingTun is a tio.Queue that records what is written to it. A queue that
|
|
||||||
// discards writes is indistinguishable from a packet that was never forwarded.
|
|
||||||
type capturingTun struct {
|
|
||||||
writes [][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *capturingTun) Read() ([]tio.Packet, error) { return nil, io.EOF }
|
|
||||||
func (c *capturingTun) Close() error { return nil }
|
|
||||||
|
|
||||||
func (c *capturingTun) Write(b []byte) (int, error) {
|
|
||||||
c.writes = append(c.writes, append([]byte(nil), b...))
|
|
||||||
return len(b), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func newSelfForwardInterface(myAddrs ...netip.Addr) (*Interface, *capturingTun) {
|
|
||||||
vpnAddrs := &bart.Lite{}
|
|
||||||
for _, a := range myAddrs {
|
|
||||||
vpnAddrs.Insert(netip.PrefixFrom(a, a.BitLen()))
|
|
||||||
}
|
|
||||||
|
|
||||||
tun := &capturingTun{}
|
|
||||||
return &Interface{
|
|
||||||
l: test.NewLogger(),
|
|
||||||
myVpnAddrsTable: vpnAddrs,
|
|
||||||
myBroadcastAddrsTable: &bart.Lite{},
|
|
||||||
queues: []tio.Queue{tun},
|
|
||||||
}, tun
|
|
||||||
}
|
|
||||||
|
|
||||||
func consumeInside(f *Interface, packet []byte) {
|
|
||||||
f.consumeInsidePacket(tio.Packet{Bytes: packet}, &firewall.ParsedPacket{}, make([]byte, 12), nil, make([]byte, mtu), 0, nil)
|
|
||||||
}
|
|
||||||
|
|
||||||
// l4Proto describes one upper-layer header for these tests: its IP next-header
|
|
||||||
// value, where its checksum field sits within the header, and how to build a
|
|
||||||
// minimal instance of it.
|
|
||||||
type l4Proto struct {
|
|
||||||
name string
|
|
||||||
nextHdr uint8
|
|
||||||
cksumAt int
|
|
||||||
build func() []byte
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
|
||||||
tcpSyn = l4Proto{"tcp", iputil.IPProtocolTCP, 16, func() []byte {
|
|
||||||
h := make([]byte, 20)
|
|
||||||
binary.BigEndian.PutUint16(h[0:2], 49152)
|
|
||||||
binary.BigEndian.PutUint16(h[2:4], 443)
|
|
||||||
binary.BigEndian.PutUint32(h[4:8], 0x11223344) // sequence
|
|
||||||
h[12] = 5 << 4 // data offset, no options
|
|
||||||
h[13] = 0x02 // SYN
|
|
||||||
binary.BigEndian.PutUint16(h[14:16], 65535) // window
|
|
||||||
return h
|
|
||||||
}}
|
|
||||||
|
|
||||||
udpDatagram = l4Proto{"udp", iputil.IPProtocolUDP, 6, func() []byte {
|
|
||||||
h := make([]byte, 8+4)
|
|
||||||
binary.BigEndian.PutUint16(h[0:2], 49152)
|
|
||||||
binary.BigEndian.PutUint16(h[2:4], 53)
|
|
||||||
binary.BigEndian.PutUint16(h[4:6], uint16(len(h)))
|
|
||||||
copy(h[8:], "ping")
|
|
||||||
return h
|
|
||||||
}}
|
|
||||||
|
|
||||||
icmpEcho = l4Proto{"icmp", iputil.IPProtocolICMP, 2, func() []byte { return echoRequest(8) }}
|
|
||||||
icmpv6Echo = l4Proto{"icmpv6", iputil.IPProtocolICMPv6, 2, func() []byte { return echoRequest(128) }}
|
|
||||||
)
|
|
||||||
|
|
||||||
// echoRequest builds an echo request body. The type differs between ICMP and
|
|
||||||
// ICMPv6, the rest of the header does not.
|
|
||||||
func echoRequest(typ uint8) []byte {
|
|
||||||
h := make([]byte, 8)
|
|
||||||
h[0] = typ
|
|
||||||
binary.BigEndian.PutUint16(h[4:6], 0xbeef) // identifier
|
|
||||||
binary.BigEndian.PutUint16(h[6:8], 1) // sequence
|
|
||||||
return h
|
|
||||||
}
|
|
||||||
|
|
||||||
func buildIPv6(src, dst netip.Addr, p l4Proto) []byte {
|
|
||||||
l4 := p.build()
|
|
||||||
pkt := make([]byte, ipv6HeaderLen+len(l4))
|
|
||||||
pkt[0] = 0x60
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(len(l4)))
|
|
||||||
pkt[6] = p.nextHdr
|
|
||||||
pkt[7] = 64
|
|
||||||
copy(pkt[8:24], src.AsSlice())
|
|
||||||
copy(pkt[24:40], dst.AsSlice())
|
|
||||||
copy(pkt[ipv6HeaderLen:], l4)
|
|
||||||
if l4 := pkt[ipv6HeaderLen:]; p.nextHdr == iputil.IPProtocolTCP || p.nextHdr == iputil.IPProtocolUDP {
|
|
||||||
sum := ipv6PseudoheaderSum(src, dst, uint32(p.nextHdr), uint32(len(l4)))
|
|
||||||
binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum)))
|
|
||||||
}
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
func buildIPv4(src, dst netip.Addr, p l4Proto) []byte {
|
|
||||||
l4 := p.build()
|
|
||||||
pkt := make([]byte, ipv4HeaderLen+len(l4))
|
|
||||||
pkt[0] = 0x45
|
|
||||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt)))
|
|
||||||
pkt[8] = 64
|
|
||||||
pkt[9] = p.nextHdr
|
|
||||||
copy(pkt[12:16], src.AsSlice())
|
|
||||||
copy(pkt[16:20], dst.AsSlice())
|
|
||||||
copy(pkt[ipv4HeaderLen:], l4)
|
|
||||||
if l4 := pkt[ipv4HeaderLen:]; p.nextHdr == iputil.IPProtocolTCP || p.nextHdr == iputil.IPProtocolUDP {
|
|
||||||
sum := sumBytes(pkt[12:20], uint32(p.nextHdr)+uint32(len(l4)))
|
|
||||||
binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum)))
|
|
||||||
}
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
// ipv6PseudoheaderSum is the RFC 2460 section 8.1 pseudo-header sum: source,
|
|
||||||
// destination, a 32 bit upper-layer packet length and a 32 bit zero-padded next
|
|
||||||
// header. Kept local to the test so these assertions do not check nebula's
|
|
||||||
// checksum code against itself.
|
|
||||||
func ipv6PseudoheaderSum(src, dst netip.Addr, nextHeader, length uint32) uint32 {
|
|
||||||
var csum uint32
|
|
||||||
s, d := src.AsSlice(), dst.AsSlice()
|
|
||||||
for i := 0; i < 16; i += 2 {
|
|
||||||
csum += uint32(s[i])<<8 | uint32(s[i+1])
|
|
||||||
csum += uint32(d[i])<<8 | uint32(d[i+1])
|
|
||||||
}
|
|
||||||
return csum + length + nextHeader
|
|
||||||
}
|
|
||||||
|
|
||||||
func sumBytes(b []byte, csum uint32) uint32 {
|
|
||||||
for i := 0; i+1 < len(b); i += 2 {
|
|
||||||
csum += uint32(b[i])<<8 | uint32(b[i+1])
|
|
||||||
}
|
|
||||||
if len(b)%2 == 1 {
|
|
||||||
csum += uint32(b[len(b)-1]) << 8
|
|
||||||
}
|
|
||||||
return csum
|
|
||||||
}
|
|
||||||
|
|
||||||
func fold(csum uint32) uint16 {
|
|
||||||
for csum > 0xffff {
|
|
||||||
csum = (csum >> 16) + (csum & 0xffff)
|
|
||||||
}
|
|
||||||
return uint16(csum)
|
|
||||||
}
|
|
||||||
|
|
||||||
// l4ChecksumValid6 verifies an IPv6 upper-layer checksum the way a receiver
|
|
||||||
// does: the pseudo-header plus the whole upper-layer segment, checksum field
|
|
||||||
// included, folds to 0xffff. The next header field is the upper-layer protocol
|
|
||||||
// only while there are no extension headers, which is all this file builds.
|
|
||||||
func l4ChecksumValid6(pkt []byte) bool {
|
|
||||||
src, _ := netip.AddrFromSlice(pkt[8:24])
|
|
||||||
dst, _ := netip.AddrFromSlice(pkt[24:40])
|
|
||||||
l4 := pkt[ipv6HeaderLen:]
|
|
||||||
return fold(sumBytes(l4, ipv6PseudoheaderSum(src, dst, uint32(pkt[6]), uint32(len(l4))))) == 0xffff
|
|
||||||
}
|
|
||||||
|
|
||||||
// l4ChecksumValid4 is the IPv4 counterpart: the RFC 793/768 pseudo-header is
|
|
||||||
// source, destination, a zero byte, the protocol and the upper-layer length.
|
|
||||||
func l4ChecksumValid4(pkt []byte) bool {
|
|
||||||
ihl := int(pkt[0]&0x0f) << 2
|
|
||||||
l4 := pkt[ihl:]
|
|
||||||
return fold(sumBytes(l4, sumBytes(pkt[12:20], uint32(pkt[9])+uint32(len(l4))))) == 0xffff
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestConsumeInsidePacketSelfTraffic covers the self-addressed branch of
|
|
||||||
// consumeInsidePacket, taken where immediatelyForwardToSelf is set (see
|
|
||||||
// inside_bsd.go): the packet goes straight back to the tun, ahead of the
|
|
||||||
// firewall and the handshake.
|
|
||||||
func TestConsumeInsidePacketSelfTraffic(t *testing.T) {
|
|
||||||
v4 := netip.MustParseAddr("100.100.1.42")
|
|
||||||
v6 := netip.MustParseAddr("fd00::42")
|
|
||||||
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
addr netip.Addr
|
|
||||||
pkt []byte
|
|
||||||
}{
|
|
||||||
{"ipv4/tcp", v4, buildIPv4(v4, v4, tcpSyn)},
|
|
||||||
{"ipv4/udp", v4, buildIPv4(v4, v4, udpDatagram)},
|
|
||||||
{"ipv4/icmp", v4, buildIPv4(v4, v4, icmpEcho)},
|
|
||||||
{"ipv6/tcp", v6, buildIPv6(v6, v6, tcpSyn)},
|
|
||||||
{"ipv6/udp", v6, buildIPv6(v6, v6, udpDatagram)},
|
|
||||||
{"ipv6/icmpv6", v6, buildIPv6(v6, v6, icmpv6Echo)},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
f, tun := newSelfForwardInterface(tt.addr)
|
|
||||||
// consumeInsidePacket writes through the slice it is handed, so a
|
|
||||||
// packet that arrived with a valid checksum must come back out of
|
|
||||||
// bytes taken before the call, unchanged.
|
|
||||||
want := append([]byte(nil), tt.pkt...)
|
|
||||||
consumeInside(f, tt.pkt)
|
|
||||||
|
|
||||||
if immediatelyForwardToSelf {
|
|
||||||
require.Len(t, tun.writes, 1)
|
|
||||||
assert.Equal(t, want, tun.writes[0])
|
|
||||||
} else {
|
|
||||||
assert.Empty(t, tun.writes, "self traffic reaches the tun over loopback here and must be dropped")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestConsumeInsidePacketSelfTrafficChecksum shows that the self-forward
|
|
||||||
// returns the bytes it was handed, so a packet that arrived with a wrong
|
|
||||||
// upper-layer checksum is written back with that same wrong checksum and the
|
|
||||||
// kernel drops it on re-entry.
|
|
||||||
//
|
|
||||||
// This is how a macOS host loses TCP and UDP to its own IPv6 overlay address:
|
|
||||||
// the kernel writes only the pseudo-header sum into the checksum field and
|
|
||||||
// defers completion to hardware offload, state that does not survive the
|
|
||||||
// crossing into userspace. Which kernels do this, for which protocols and IP
|
|
||||||
// versions, is a property of the kernel and belongs to a test against a live
|
|
||||||
// one; here the checksum is simply wrong, and the forward must make it right.
|
|
||||||
func TestConsumeInsidePacketSelfTrafficChecksum(t *testing.T) {
|
|
||||||
if !immediatelyForwardToSelf {
|
|
||||||
t.Skip("self traffic never reaches the tun on this platform")
|
|
||||||
}
|
|
||||||
versions := []struct {
|
|
||||||
name string
|
|
||||||
addr netip.Addr
|
|
||||||
build func(src, dst netip.Addr, p l4Proto) []byte
|
|
||||||
l4At int
|
|
||||||
valid func(pkt []byte) bool
|
|
||||||
}{
|
|
||||||
{"v4", netip.MustParseAddr("100.100.1.42"), buildIPv4, ipv4HeaderLen, l4ChecksumValid4},
|
|
||||||
{"v6", netip.MustParseAddr("fd00::42"), buildIPv6, ipv6HeaderLen, l4ChecksumValid6},
|
|
||||||
}
|
|
||||||
for _, v := range versions {
|
|
||||||
for _, p := range []l4Proto{tcpSyn, udpDatagram} {
|
|
||||||
t.Run(v.name+"/"+p.name, func(t *testing.T) {
|
|
||||||
pkt := v.build(v.addr, v.addr, p)
|
|
||||||
binary.BigEndian.PutUint16(pkt[v.l4At+p.cksumAt:], 0x1234)
|
|
||||||
require.False(t, v.valid(pkt), "the packet under test must start with a wrong checksum")
|
|
||||||
f, tun := newSelfForwardInterface(v.addr)
|
|
||||||
consumeInside(f, pkt)
|
|
||||||
require.Len(t, tun.writes, 1)
|
|
||||||
assert.True(t, v.valid(tun.writes[0]),
|
|
||||||
"a forwarded %s packet must carry a valid checksum, got 0x%04x",
|
|
||||||
p.name, binary.BigEndian.Uint16(tun.writes[0][v.l4At+p.cksumAt:]))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+46
-171
@@ -2,12 +2,11 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"crypto/fips140"
|
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"runtime"
|
|
||||||
"slices"
|
"slices"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
@@ -15,15 +14,12 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/util"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/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"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -53,19 +49,7 @@ type InterfaceConfig struct {
|
|||||||
reQueryWait time.Duration
|
reQueryWait time.Duration
|
||||||
|
|
||||||
ConntrackCacheTimeout time.Duration
|
ConntrackCacheTimeout time.Duration
|
||||||
|
l *slog.Logger
|
||||||
// CpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
|
|
||||||
// should pin to. Queue i pins to CpuAffinity[i % len(CpuAffinity)] —
|
|
||||||
// shorter lists than `routines` cycle. Empty list keeps the default
|
|
||||||
// pin-to-(i % NumCPU) behavior. Only consulted when PinThreads is true.
|
|
||||||
CpuAffinity []int
|
|
||||||
// PinThreads controls whether each TUN reader OS thread is pinned to a
|
|
||||||
// single CPU (via tun.pin_threads, default true). Pinning keeps each
|
|
||||||
// goroutine's sendmmsg on one XPS-selected NIC TX ring so per-flow
|
|
||||||
// packets stay ordered on the wire.
|
|
||||||
PinThreads bool
|
|
||||||
|
|
||||||
l *slog.Logger
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type Interface struct {
|
type Interface struct {
|
||||||
@@ -89,16 +73,7 @@ type Interface struct {
|
|||||||
routines int
|
routines int
|
||||||
disconnectInvalid atomic.Bool
|
disconnectInvalid atomic.Bool
|
||||||
closed atomic.Bool
|
closed atomic.Bool
|
||||||
// cpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
|
relayManager *relayManager
|
||||||
// should pin to. Queue i pins to cpuAffinity[i % len(cpuAffinity)].
|
|
||||||
// Empty falls back to the default pin-to-(allowed CPU) behavior.
|
|
||||||
// Only consulted when pinThreads is true.
|
|
||||||
cpuAffinity []int
|
|
||||||
// pinThreads controls whether listenIn pins each TUN reader OS thread to
|
|
||||||
// a CPU at all (tun.pin_threads, default true). When false, threads are
|
|
||||||
// left free to migrate as on stock nebula.
|
|
||||||
pinThreads bool
|
|
||||||
relayManager *relayManager
|
|
||||||
|
|
||||||
tryPromoteEvery atomic.Uint32
|
tryPromoteEvery atomic.Uint32
|
||||||
reQueryEvery atomic.Uint32
|
reQueryEvery atomic.Uint32
|
||||||
@@ -107,22 +82,16 @@ type Interface struct {
|
|||||||
sendRecvErrorConfig recvErrorConfig
|
sendRecvErrorConfig recvErrorConfig
|
||||||
acceptRecvErrorConfig recvErrorConfig
|
acceptRecvErrorConfig recvErrorConfig
|
||||||
|
|
||||||
// Bumped on every udp rebind, tunnels compare it to decide they need a punch from the far side
|
// rebindCount is used to decide if an active tunnel should trigger a punch notification through a lighthouse
|
||||||
rebindEpoch atomic.Uint32
|
rebindCount int8
|
||||||
version string
|
version string
|
||||||
|
|
||||||
conntrackCacheTimeout time.Duration
|
conntrackCacheTimeout time.Duration
|
||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
writers []udp.Conn
|
writers []udp.Conn
|
||||||
queues []tio.Queue
|
readers []io.ReadWriteCloser
|
||||||
// batchers is one per tun queue, wrapping queues[i]. readOutsidePackets
|
wg sync.WaitGroup
|
||||||
// commits plaintext into the batcher; the plaintext is decrypted
|
|
||||||
// in place inside the UDP receive buffers, so listenOut must call Flush
|
|
||||||
// at the end of each UDP recvmmsg batch, before those buffers are
|
|
||||||
// reused (every udp.Conn ListenOut guarantees that ordering).
|
|
||||||
batchers []*batch.MultiCoalescer
|
|
||||||
wg sync.WaitGroup
|
|
||||||
|
|
||||||
// fatalErr holds the first unexpected reader error that caused shutdown.
|
// fatalErr holds the first unexpected reader error that caused shutdown.
|
||||||
// nil means "no fatal error" (yet)
|
// nil means "no fatal error" (yet)
|
||||||
@@ -133,13 +102,18 @@ type Interface struct {
|
|||||||
metricHandshakes metrics.Histogram
|
metricHandshakes metrics.Histogram
|
||||||
messageMetrics *MessageMetrics
|
messageMetrics *MessageMetrics
|
||||||
cachedPacketMetrics *cachedPacketMetrics
|
cachedPacketMetrics *cachedPacketMetrics
|
||||||
metricTxDropped metrics.Counter
|
|
||||||
|
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
type EncWriter interface {
|
type EncWriter interface {
|
||||||
SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int)
|
SendVia(via *HostInfo,
|
||||||
|
relay *Relay,
|
||||||
|
ad,
|
||||||
|
nb,
|
||||||
|
out []byte,
|
||||||
|
nocopy bool,
|
||||||
|
)
|
||||||
SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p, nb, out []byte)
|
SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p, nb, out []byte)
|
||||||
SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte)
|
SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte)
|
||||||
Handshake(vpnAddr netip.Addr)
|
Handshake(vpnAddr netip.Addr)
|
||||||
@@ -198,10 +172,6 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
return nil, errors.New("no connection manager")
|
return nil, errors.New("no connection manager")
|
||||||
}
|
}
|
||||||
|
|
||||||
if c.routines <= 1 {
|
|
||||||
c.PinThreads = false //pinning is not useful unless there's more than one tun reader
|
|
||||||
}
|
|
||||||
|
|
||||||
cs := c.pki.getCertState()
|
cs := c.pki.getCertState()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
@@ -219,7 +189,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
routines: c.routines,
|
routines: c.routines,
|
||||||
version: c.version,
|
version: c.version,
|
||||||
writers: make([]udp.Conn, c.routines),
|
writers: make([]udp.Conn, c.routines),
|
||||||
batchers: make([]*batch.MultiCoalescer, c.routines),
|
readers: make([]io.ReadWriteCloser, c.routines),
|
||||||
myVpnNetworks: cs.myVpnNetworks,
|
myVpnNetworks: cs.myVpnNetworks,
|
||||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||||
myVpnAddrs: cs.myVpnAddrs,
|
myVpnAddrs: cs.myVpnAddrs,
|
||||||
@@ -228,11 +198,8 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
relayManager: c.relayManager,
|
relayManager: c.relayManager,
|
||||||
connectionManager: c.connectionManager,
|
connectionManager: c.connectionManager,
|
||||||
conntrackCacheTimeout: c.ConntrackCacheTimeout,
|
conntrackCacheTimeout: c.ConntrackCacheTimeout,
|
||||||
cpuAffinity: c.CpuAffinity,
|
|
||||||
pinThreads: c.PinThreads,
|
|
||||||
|
|
||||||
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
||||||
metricTxDropped: metrics.GetOrRegisterCounter("udp.tx.dropped", nil),
|
|
||||||
messageMetrics: c.MessageMetrics,
|
messageMetrics: c.MessageMetrics,
|
||||||
cachedPacketMetrics: &cachedPacketMetrics{
|
cachedPacketMetrics: &cachedPacketMetrics{
|
||||||
sent: metrics.GetOrRegisterCounter("hostinfo.cached_packets.sent", nil),
|
sent: metrics.GetOrRegisterCounter("hostinfo.cached_packets.sent", nil),
|
||||||
@@ -271,41 +238,27 @@ func (f *Interface) activate() error {
|
|||||||
"build", f.version,
|
"build", f.version,
|
||||||
"udpAddr", addr,
|
"udpAddr", addr,
|
||||||
"boringcrypto", boringEnabled(),
|
"boringcrypto", boringEnabled(),
|
||||||
"fips140Version", fips140.Version(),
|
|
||||||
"fips140Enabled", fips140.Enabled(),
|
|
||||||
"fips140Enforced", fips140.Enforced(),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if f.routines > 1 && !f.outside.SupportsMultipleReaders() {
|
if f.routines > 1 {
|
||||||
f.routines = 1
|
if !f.inside.SupportsMultiqueue() || !f.outside.SupportsMultipleReaders() {
|
||||||
f.l.Warn("multiple udp readers are not supported on this platform, falling back to a single routine")
|
f.routines = 1
|
||||||
|
f.l.Warn("routines is not supported on this platform, falling back to a single routine")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Prepare the tun queues. A device that can't open that many hands back
|
|
||||||
// fewer (a single queue on platforms without multiqueue support) and we
|
|
||||||
// size the reader routines to what we actually got.
|
|
||||||
queues, err := f.inside.Queues(f.routines)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if len(queues) < f.routines {
|
|
||||||
// TODO: this clamp is only safe because it is unreachable when the
|
|
||||||
// udp side has multiple readers (linux Queues opens exactly n or
|
|
||||||
// errors; every other platform already clamped routines to 1 above).
|
|
||||||
// If a platform ever returns fewer queues than routines with
|
|
||||||
// SO_REUSEPORT sockets already bound, the surplus sockets get no
|
|
||||||
// listenOut and the kernel blackholes every flow it hashes to them —
|
|
||||||
// fail loudly or close the extra sockets instead.
|
|
||||||
f.l.Warn("tun multiqueue is not supported on this platform, falling back to fewer routines",
|
|
||||||
"requested", f.routines, "opened", len(queues))
|
|
||||||
f.routines = len(queues)
|
|
||||||
}
|
|
||||||
f.queues = queues
|
|
||||||
|
|
||||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||||
|
|
||||||
for i := range f.queues {
|
// Prepare n tun queues
|
||||||
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l)
|
var reader io.ReadWriteCloser = f.inside
|
||||||
|
for i := 0; i < f.routines; i++ {
|
||||||
|
if i > 0 {
|
||||||
|
reader, err = f.inside.NewMultiQueueReader()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
f.readers[i] = reader
|
||||||
}
|
}
|
||||||
|
|
||||||
// On error the caller owns the cleanup, Control.Start cancels the service context
|
// On error the caller owns the cleanup, Control.Start cancels the service context
|
||||||
@@ -328,7 +281,7 @@ func (f *Interface) run() {
|
|||||||
// Launch n queues to read packets from tun dev
|
// Launch n queues to read packets from tun dev
|
||||||
for i := 0; i < f.routines; i++ {
|
for i := 0; i < f.routines; i++ {
|
||||||
f.wg.Go(func() {
|
f.wg.Go(func() {
|
||||||
f.listenIn(f.queues[i], i)
|
f.listenIn(f.readers[i], i)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -353,31 +306,6 @@ func (f *Interface) onFatal(err error) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
type rxContext struct {
|
|
||||||
q int
|
|
||||||
scratch []byte
|
|
||||||
// nb is a re-usable nonce buffer for decrypt calls to use
|
|
||||||
nb []byte
|
|
||||||
h *header.H
|
|
||||||
fwPacket *firewall.ParsedPacket
|
|
||||||
hostmapCache map[uint32]*HostInfo
|
|
||||||
lhh *LightHouseHandler
|
|
||||||
ctCache *firewall.ConntrackCacheTicker
|
|
||||||
}
|
|
||||||
|
|
||||||
func newRxContext(f *Interface, q int) *rxContext {
|
|
||||||
return &rxContext{
|
|
||||||
q: q,
|
|
||||||
scratch: make([]byte, mtu),
|
|
||||||
nb: make([]byte, 12, 12),
|
|
||||||
h: &header.H{},
|
|
||||||
fwPacket: &firewall.ParsedPacket{},
|
|
||||||
hostmapCache: map[uint32]*HostInfo{},
|
|
||||||
lhh: f.lightHouse.NewRequestHandler(),
|
|
||||||
ctCache: firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *Interface) listenOut(i int) {
|
func (f *Interface) listenOut(i int) {
|
||||||
var li udp.Conn
|
var li udp.Conn
|
||||||
if i > 0 {
|
if i > 0 {
|
||||||
@@ -386,20 +314,16 @@ func (f *Interface) listenOut(i int) {
|
|||||||
li = f.outside
|
li = f.outside
|
||||||
}
|
}
|
||||||
|
|
||||||
rxc := newRxContext(f, i)
|
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
|
lhh := f.lightHouse.NewRequestHandler()
|
||||||
|
plaintext := make([]byte, udp.MTU)
|
||||||
|
h := &header.H{}
|
||||||
|
fwPacket := &firewall.Packet{}
|
||||||
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
listener := func(fromUdpAddr netip.AddrPort, payload []byte) {
|
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, payload, rxc)
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
||||||
}
|
})
|
||||||
|
|
||||||
flusher := func() {
|
|
||||||
if err := f.batchers[i].Flush(); err != nil {
|
|
||||||
f.l.Error("Failed to flush tun coalescer", "error", err)
|
|
||||||
}
|
|
||||||
clear(rxc.hostmapCache)
|
|
||||||
}
|
|
||||||
|
|
||||||
err := li.ListenOut(listener, flusher)
|
|
||||||
|
|
||||||
// An error after teardown began is shutdown noise, the closed flag covers resources
|
// An error after teardown began is shutdown noise, the closed flag covers resources
|
||||||
// Close releases itself and the cancelled ctx covers ones torn down by their owners
|
// Close releases itself and the cancelled ctx covers ones torn down by their owners
|
||||||
@@ -412,42 +336,16 @@ 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) pinThisThread(i int) {
|
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
||||||
var cpu int
|
packet := make([]byte, mtu)
|
||||||
if n := len(f.cpuAffinity); n > 0 {
|
out := make([]byte, mtu)
|
||||||
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already
|
fwPacket := &firewall.Packet{}
|
||||||
// validated the entries against the allowed CPU set.
|
|
||||||
cpu = f.cpuAffinity[i%n]
|
|
||||||
} else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 {
|
|
||||||
// Default: spread queues across the CPUs we're actually allowed to
|
|
||||||
// run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so
|
|
||||||
// i % NumCPU would pick unrunnable IDs and every pin would fail.
|
|
||||||
cpu = allowed[i%len(allowed)]
|
|
||||||
} else {
|
|
||||||
cpu = i % runtime.NumCPU()
|
|
||||||
}
|
|
||||||
if err := util.PinThreadToCPU(cpu); err != nil {
|
|
||||||
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *Interface) listenIn(queue tio.Queue, i int) {
|
|
||||||
// Pinning this thread (and goroutine) to a single CPU keeps every sendmmsg from this goroutine going through the
|
|
||||||
// same TX ring on the nic, so the wire sees per-flow order. Skip entirely when tun.pin_threads is false.
|
|
||||||
if f.pinThreads {
|
|
||||||
f.pinThisThread(i)
|
|
||||||
}
|
|
||||||
|
|
||||||
rejectBuf := make([]byte, mtu)
|
|
||||||
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
|
|
||||||
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
|
|
||||||
fwPacket := &firewall.ParsedPacket{}
|
|
||||||
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 {
|
||||||
pkts, err := queue.Read()
|
n, err := reader.Read(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Same shutdown noise handling as listenOut
|
// Same shutdown noise handling as listenOut
|
||||||
if !f.closed.Load() && f.ctx.Err() == nil {
|
if !f.closed.Load() && f.ctx.Err() == nil {
|
||||||
@@ -457,35 +355,12 @@ func (f *Interface) listenIn(queue tio.Queue, i int) {
|
|||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, pkt := range pkts {
|
f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get())
|
||||||
f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get())
|
|
||||||
// Flush incrementally once a full sendmmsg batch has
|
|
||||||
// accumulated so the first packets of a deep read drain
|
|
||||||
// hit the wire while the rest are still being encrypted.
|
|
||||||
if sb.Len() >= batch.SendBatchCap {
|
|
||||||
f.flushSendBatch(sb, i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
f.flushSendBatch(sb, i)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
f.l.Debug("overlay reader is done", "reader", i)
|
f.l.Debug("overlay reader is done", "reader", i)
|
||||||
}
|
}
|
||||||
|
|
||||||
// flushSendBatch drains sb to the underlay and accounts for anything it could not deliver. A shortfall means
|
|
||||||
// specific destinations were undeliverable (a stale remote, a reject rule), which the backend logs per peer at
|
|
||||||
// debug; here it is only a counter, so one unreachable peer cannot spam a log line per batch.
|
|
||||||
func (f *Interface) flushSendBatch(sb *batch.SendBatch, q int) {
|
|
||||||
queued := sb.Len()
|
|
||||||
written, err := sb.Flush()
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to write outgoing batch", "error", err, "writer", q)
|
|
||||||
}
|
|
||||||
if dropped := queued - written; dropped > 0 {
|
|
||||||
f.metricTxDropped.Inc(int64(dropped))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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)
|
||||||
|
|||||||
@@ -1,146 +0,0 @@
|
|||||||
package iputil
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/checksum"
|
|
||||||
"golang.org/x/net/ipv4"
|
|
||||||
"golang.org/x/net/ipv6"
|
|
||||||
)
|
|
||||||
|
|
||||||
const udpHeaderLen = 8
|
|
||||||
|
|
||||||
// SetTransportChecksum recomputes the TCP or UDP checksum of an IPv4 or IPv6
|
|
||||||
// packet in place.
|
|
||||||
//
|
|
||||||
// A kernel that offloads checksums to the NIC hands a packet to a tun with the
|
|
||||||
// transport checksum unfinished: only the pseudo-header sum is in the field and
|
|
||||||
// the rest is left for hardware that a tun does not have. A packet written
|
|
||||||
// straight back to that tun is dropped on re-entry unless the checksum is
|
|
||||||
// completed first. ICMP is left alone; it arrived complete on the kernels this
|
|
||||||
// was measured against.
|
|
||||||
//
|
|
||||||
// So is any packet whose transport header cannot be located: fragments, unknown
|
|
||||||
// extension headers and truncated packets. An IPv6 fragment header is declined
|
|
||||||
// even when it carries the whole datagram (RFC 6946 atomic fragment), because
|
|
||||||
// the walk reports only that a fragment header was present.
|
|
||||||
func SetTransportChecksum(packet []byte) {
|
|
||||||
if len(packet) < 1 {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
switch int(packet[0] >> 4) {
|
|
||||||
case ipv4.Version:
|
|
||||||
setTransportChecksum4(packet)
|
|
||||||
case ipv6.Version:
|
|
||||||
setTransportChecksum6(packet)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func setTransportChecksum4(packet []byte) {
|
|
||||||
if len(packet) < ipv4.HeaderLen {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
ihl := int(packet[0]&0x0f) << 2
|
|
||||||
end := int(binary.BigEndian.Uint16(packet[2:4]))
|
|
||||||
if ihl < ipv4.HeaderLen || end < ihl || end > len(packet) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
// The checksum covers the whole datagram, which a fragment (MF set or a
|
|
||||||
// non-zero offset) does not carry.
|
|
||||||
if binary.BigEndian.Uint16(packet[6:8])&0x3fff != 0 {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
transport, ok := transportExtent(packet[ihl:end], packet[9])
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
csum := ipv4PseudoheaderChecksum(packet[12:16], packet[16:20], uint32(packet[9]), uint32(len(transport)))
|
|
||||||
writeTransportChecksum(transport, packet[9], csum)
|
|
||||||
}
|
|
||||||
|
|
||||||
func setTransportChecksum6(packet []byte) {
|
|
||||||
if len(packet) < ipv6.HeaderLen {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
end := ipv6.HeaderLen + int(binary.BigEndian.Uint16(packet[4:6]))
|
|
||||||
if end > len(packet) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// The checksum covers the whole datagram, which a fragment does not carry.
|
|
||||||
// An unknown extension header hides where the transport header starts. A
|
|
||||||
// chain longer than the walk's budget ends it early, at an offset that was
|
|
||||||
// never checked against the packet.
|
|
||||||
proto, offset, _, anyFragment, err := IPv6FindUpperProtocol(packet[:end])
|
|
||||||
if err != nil || anyFragment || offset >= end {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
transport, ok := transportExtent(packet[offset:end], proto)
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
csum := ipv6PseudoheaderChecksum(packet[8:24], packet[24:40], uint32(proto), uint32(len(transport)))
|
|
||||||
writeTransportChecksum(transport, proto, csum)
|
|
||||||
}
|
|
||||||
|
|
||||||
// transportExtent narrows a segment to the length its own header declares. UDP
|
|
||||||
// carries a Length field, and RFC 768 and RFC 8200 section 8.1 both make that
|
|
||||||
// field, not the IP payload extent, the length the pseudo-header counts and the
|
|
||||||
// checksum covers; a datagram padded out to a link's minimum frame is the usual
|
|
||||||
// way the two differ. TCP has no such field, so its segment runs to the end of
|
|
||||||
// the IP payload. A Length that overruns the bytes IP delivered describes a
|
|
||||||
// datagram that is not there.
|
|
||||||
func transportExtent(transport []byte, proto uint8) ([]byte, bool) {
|
|
||||||
if proto != IPProtocolUDP {
|
|
||||||
return transport, true
|
|
||||||
}
|
|
||||||
if len(transport) < udpHeaderLen {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
ulen := int(binary.BigEndian.Uint16(transport[4:6]))
|
|
||||||
if ulen < udpHeaderLen || ulen > len(transport) {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
return transport[:ulen], true
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeTransportChecksum stores the checksum of transport, taken over the
|
|
||||||
// pseudo-header sum csum, in the header's checksum field. A UDP checksum that
|
|
||||||
// computes to zero goes on the wire as 0xffff: zero means no checksum was
|
|
||||||
// computed (RFC 768), and over IPv6 the checksum is mandatory (RFC 8200
|
|
||||||
// section 8.1).
|
|
||||||
func writeTransportChecksum(transport []byte, proto uint8, csum uint32) {
|
|
||||||
var at, minLen int
|
|
||||||
switch proto {
|
|
||||||
case IPProtocolTCP:
|
|
||||||
at, minLen = 16, 20
|
|
||||||
case IPProtocolUDP:
|
|
||||||
at, minLen = 6, udpHeaderLen
|
|
||||||
default:
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if len(transport) < minLen {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
transport[at], transport[at+1] = 0, 0
|
|
||||||
sum := ^checksum.Checksum(transport, fold(csum))
|
|
||||||
if sum == 0 && proto == IPProtocolUDP {
|
|
||||||
sum = 0xffff
|
|
||||||
}
|
|
||||||
binary.BigEndian.PutUint16(transport[at:], sum)
|
|
||||||
}
|
|
||||||
|
|
||||||
// fold reduces a pseudo-header sum to the 16 bit seed Checksum takes. Carrying
|
|
||||||
// the high half back into the low half is what keeps the reduction lossless, so
|
|
||||||
// the seed sums exactly as the wider value would; 0xffff is its fixed point.
|
|
||||||
// Every term of that sum comes from a 16 bit field, so it stays far below the
|
|
||||||
// width at which the accumulator would wrap.
|
|
||||||
func fold(csum uint32) uint16 {
|
|
||||||
for csum > 0xffff {
|
|
||||||
csum = (csum >> 16) + (csum & 0xffff)
|
|
||||||
}
|
|
||||||
return uint16(csum)
|
|
||||||
}
|
|
||||||
@@ -1,242 +0,0 @@
|
|||||||
package iputil
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"net"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/google/gopacket"
|
|
||||||
"github.com/google/gopacket/layers"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"golang.org/x/net/ipv6"
|
|
||||||
)
|
|
||||||
|
|
||||||
// serialize builds a packet with gopacket, whose checksums are computed
|
|
||||||
// independently of this package.
|
|
||||||
func serialize(t *testing.T, ls ...gopacket.SerializableLayer) []byte {
|
|
||||||
buf := gopacket.NewSerializeBuffer()
|
|
||||||
require.NoError(t, gopacket.SerializeLayers(buf, gopacket.SerializeOptions{FixLengths: true, ComputeChecksums: true}, ls...))
|
|
||||||
return append([]byte(nil), buf.Bytes()...)
|
|
||||||
}
|
|
||||||
|
|
||||||
// withExtensionHeader inserts an 8 byte IPv6 extension header of the given
|
|
||||||
// type between the IPv6 header and its payload. The transport checksum does not
|
|
||||||
// change: the pseudo-header counts only upper-layer bytes.
|
|
||||||
func withExtensionHeader(pkt []byte, typ layers.IPProtocol, hdr [8]byte) []byte {
|
|
||||||
hdr[0] = pkt[6]
|
|
||||||
out := make([]byte, 0, len(pkt)+8)
|
|
||||||
out = append(out, pkt[:40]...)
|
|
||||||
out = append(out, hdr[:]...)
|
|
||||||
out = append(out, pkt[40:]...)
|
|
||||||
out[6] = byte(typ)
|
|
||||||
binary.BigEndian.PutUint16(out[4:6], binary.BigEndian.Uint16(pkt[4:6])+8)
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
// truncate copies the first n bytes into a buffer of exactly that capacity, so
|
|
||||||
// a read past the length panics instead of quietly succeeding.
|
|
||||||
func truncate(pkt []byte, n int) []byte {
|
|
||||||
out := make([]byte, n)
|
|
||||||
copy(out, pkt)
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
// extChain builds an IPv6 packet fronted by n Destination Options headers. Each
|
|
||||||
// points at another one, so the walk spends its whole budget without reaching a
|
|
||||||
// transport header. lastExtLen inflates the final header's declared length,
|
|
||||||
// which is how the walk ends up past the end of the packet.
|
|
||||||
func extChain(n int, lastExtLen byte) []byte {
|
|
||||||
pkt := make([]byte, ipv6.HeaderLen)
|
|
||||||
pkt[0], pkt[6], pkt[7] = 0x60, 60, 64
|
|
||||||
for i := range n {
|
|
||||||
h := make([]byte, 8)
|
|
||||||
h[0] = 60
|
|
||||||
if i == n-1 {
|
|
||||||
h[1] = lastExtLen
|
|
||||||
}
|
|
||||||
pkt = append(pkt, h...)
|
|
||||||
}
|
|
||||||
pkt = append(pkt, make([]byte, 20)...)
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(len(pkt)-ipv6.HeaderLen))
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSetTransportChecksum(t *testing.T) {
|
|
||||||
// Source and destination differ so that a pseudo-header built from the wrong
|
|
||||||
// one, or from the two swapped, does not land on the same checksum anyway.
|
|
||||||
v4 := func(proto layers.IPProtocol) *layers.IPv4 {
|
|
||||||
return &layers.IPv4{Version: 4, TTL: 64, Id: 0x1234, Protocol: proto, SrcIP: net.IPv4(192, 0, 2, 1).To4(), DstIP: net.IPv4(198, 51, 100, 2).To4()}
|
|
||||||
}
|
|
||||||
v6 := func(proto layers.IPProtocol) *layers.IPv6 {
|
|
||||||
return &layers.IPv6{Version: 6, HopLimit: 64, NextHeader: proto, SrcIP: net.ParseIP("2001:db8::1"), DstIP: net.ParseIP("2001:db8:1::2")}
|
|
||||||
}
|
|
||||||
tcp := func(ip gopacket.NetworkLayer) *layers.TCP {
|
|
||||||
l := &layers.TCP{SrcPort: 49152, DstPort: 443, SYN: true, Window: 65535}
|
|
||||||
require.NoError(t, l.SetNetworkLayerForChecksum(ip))
|
|
||||||
return l
|
|
||||||
}
|
|
||||||
udp := func(ip gopacket.NetworkLayer) *layers.UDP {
|
|
||||||
l := &layers.UDP{SrcPort: 49152, DstPort: 53}
|
|
||||||
require.NoError(t, l.SetNetworkLayerForChecksum(ip))
|
|
||||||
return l
|
|
||||||
}
|
|
||||||
payload := gopacket.Payload("self")
|
|
||||||
nop := layers.IPv4Option{OptionType: 1, OptionLength: 1}
|
|
||||||
|
|
||||||
ip4tcp := v4(layers.IPProtocolTCP)
|
|
||||||
ip4opts := v4(layers.IPProtocolTCP)
|
|
||||||
ip4opts.Options = []layers.IPv4Option{nop, nop, nop, nop}
|
|
||||||
ip4udp := v4(layers.IPProtocolUDP)
|
|
||||||
ip6tcp := v6(layers.IPProtocolTCP)
|
|
||||||
ip6udp := v6(layers.IPProtocolUDP)
|
|
||||||
hopByHop := [8]byte{0, 0, 1, 4} // next header, length 0, PadN of 4
|
|
||||||
|
|
||||||
// Bytes past the length the IP header declares are not part of the
|
|
||||||
// datagram and must not be summed.
|
|
||||||
trailing4 := append(serialize(t, ip4tcp, tcp(ip4tcp), payload), []byte("trailing")...)
|
|
||||||
trailing6 := append(serialize(t, ip6tcp, tcp(ip6tcp), payload), []byte("trailing")...)
|
|
||||||
|
|
||||||
// A datagram padded out past the length UDP declares: the pseudo-header
|
|
||||||
// counts the UDP Length field, so the checksum is the unpadded one.
|
|
||||||
padded4 := append(serialize(t, ip4udp, udp(ip4udp), payload), []byte("pad!")...)
|
|
||||||
binary.BigEndian.PutUint16(padded4[2:4], uint16(len(padded4)))
|
|
||||||
padded6 := append(serialize(t, ip6udp, udp(ip6udp), payload), []byte("pad!")...)
|
|
||||||
binary.BigEndian.PutUint16(padded6[4:6], uint16(len(padded6)-ipv6.HeaderLen))
|
|
||||||
|
|
||||||
// Corrupting the checksum and asking for it back must yield gopacket's
|
|
||||||
// packet, byte for byte.
|
|
||||||
recomputed := []struct {
|
|
||||||
name string
|
|
||||||
pkt []byte
|
|
||||||
cksum int
|
|
||||||
}{
|
|
||||||
{"v4 tcp", serialize(t, ip4tcp, tcp(ip4tcp), payload), 20 + 16},
|
|
||||||
{"v4 tcp with ip options", serialize(t, ip4opts, tcp(ip4opts), payload), 24 + 16},
|
|
||||||
{"v4 udp", serialize(t, ip4udp, udp(ip4udp), payload), 20 + 6},
|
|
||||||
{"v4 tcp header only", serialize(t, ip4tcp, tcp(ip4tcp)), 20 + 16},
|
|
||||||
{"v4 udp header only", serialize(t, ip4udp, udp(ip4udp)), 20 + 6},
|
|
||||||
{"v6 tcp", serialize(t, ip6tcp, tcp(ip6tcp), payload), 40 + 16},
|
|
||||||
{"v6 udp", serialize(t, ip6udp, udp(ip6udp), payload), 40 + 6},
|
|
||||||
{"v6 udp header only", serialize(t, ip6udp, udp(ip6udp)), 40 + 6},
|
|
||||||
{"v6 tcp behind hop-by-hop", withExtensionHeader(serialize(t, ip6tcp, tcp(ip6tcp), payload), layers.IPProtocolIPv6HopByHop, hopByHop), 48 + 16},
|
|
||||||
{"v4 tcp with bytes past the total length", trailing4, 20 + 16},
|
|
||||||
{"v6 tcp with bytes past the payload length", trailing6, 40 + 16},
|
|
||||||
{"v4 udp padded past its declared length", padded4, 20 + 6},
|
|
||||||
{"v6 udp padded past its declared length", padded6, 40 + 6},
|
|
||||||
}
|
|
||||||
for _, tt := range recomputed {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
got := append([]byte(nil), tt.pkt...)
|
|
||||||
binary.BigEndian.PutUint16(got[tt.cksum:], 0x1234)
|
|
||||||
require.NotEqual(t, tt.pkt, got)
|
|
||||||
SetTransportChecksum(got)
|
|
||||||
assert.Equal(t, tt.pkt, got)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
ip4frag := v4(layers.IPProtocolTCP)
|
|
||||||
ip4frag.Flags = layers.IPv4MoreFragments
|
|
||||||
ip4later := v4(layers.IPProtocolTCP)
|
|
||||||
ip4later.FragOffset = 1
|
|
||||||
ip4icmp := v4(layers.IPProtocolICMPv4)
|
|
||||||
|
|
||||||
badIHL := serialize(t, ip4tcp, tcp(ip4tcp), payload)
|
|
||||||
badIHL[0] = 0x44 // header length 16, shorter than an ipv4 header
|
|
||||||
shortTotalLen := serialize(t, ip4tcp, tcp(ip4tcp), payload)
|
|
||||||
binary.BigEndian.PutUint16(shortTotalLen[2:4], 10) // shorter than the header it introduces
|
|
||||||
cutTCP := serialize(t, ip4tcp, tcp(ip4tcp), payload)
|
|
||||||
binary.BigEndian.PutUint16(cutTCP[2:4], 20+19) // one byte short of a tcp header
|
|
||||||
cutTCP = truncate(cutTCP, 20+19)
|
|
||||||
cutUDP := serialize(t, ip4udp, udp(ip4udp), payload)
|
|
||||||
binary.BigEndian.PutUint16(cutUDP[2:4], 20+7) // one byte short of a udp header
|
|
||||||
cutUDP = truncate(cutUDP, 20+7)
|
|
||||||
// Two bytes short, so a transport header survives whole and the minimum
|
|
||||||
// length check cannot stand in for the bounds check.
|
|
||||||
cutV6 := truncate(serialize(t, ip6tcp, tcp(ip6tcp), payload), 62)
|
|
||||||
fragment := [8]byte{0, 0, 0, 1, 0, 0, 0, 1} // next header, reserved, offset 0 with M set, id
|
|
||||||
overrun4 := serialize(t, ip4udp, udp(ip4udp), payload)
|
|
||||||
binary.BigEndian.PutUint16(overrun4[24:26], uint16(len(overrun4)-20+1)) // one byte past what ip delivered
|
|
||||||
overrun6 := serialize(t, ip6udp, udp(ip6udp), payload)
|
|
||||||
binary.BigEndian.PutUint16(overrun6[44:46], uint16(len(overrun6)-ipv6.HeaderLen+1))
|
|
||||||
shortUDPLen := serialize(t, ip4udp, udp(ip4udp), payload)
|
|
||||||
binary.BigEndian.PutUint16(shortUDPLen[24:26], 7) // shorter than the header it counts
|
|
||||||
|
|
||||||
// Where the checksum cannot be completed the packet is left as it came.
|
|
||||||
untouched := []struct {
|
|
||||||
name string
|
|
||||||
pkt []byte
|
|
||||||
cksum int
|
|
||||||
}{
|
|
||||||
{"v4 first fragment", serialize(t, ip4frag, tcp(ip4frag), payload), 20 + 16},
|
|
||||||
{"v4 later fragment", serialize(t, ip4later, tcp(ip4later), payload), 20 + 16},
|
|
||||||
{"v4 icmp", serialize(t, ip4icmp, &layers.ICMPv4{TypeCode: layers.CreateICMPv4TypeCode(8, 0), Id: 1, Seq: 1}, payload), 20 + 2},
|
|
||||||
{"v4 header length below the minimum", badIHL, 20 + 16},
|
|
||||||
{"v4 total length below the header length", shortTotalLen, 20 + 16},
|
|
||||||
{"v4 truncated below its total length", truncate(serialize(t, ip4tcp, tcp(ip4tcp), payload), 30), -1},
|
|
||||||
{"v4 tcp header cut short", cutTCP, 20 + 16},
|
|
||||||
{"v4 udp header cut short", cutUDP, -1},
|
|
||||||
{"v6 fragment", withExtensionHeader(serialize(t, ip6tcp, tcp(ip6tcp), payload), layers.IPProtocolIPv6Fragment, fragment), 48 + 16},
|
|
||||||
{"v6 truncated below its payload length", truncate(serialize(t, ip6tcp, tcp(ip6tcp), payload), 50), -1},
|
|
||||||
{"v6 truncated with a whole transport header still present", cutV6, 40 + 16},
|
|
||||||
{"v6 extension header chain longer than the walk", extChain(9, 0), 112 + 16},
|
|
||||||
{"v6 extension header chain running past the packet", extChain(8, 255), 104 + 16},
|
|
||||||
{"v4 udp length past the end of the datagram", overrun4, 20 + 6},
|
|
||||||
{"v6 udp length past the end of the datagram", overrun6, 40 + 6},
|
|
||||||
{"v4 udp length below a udp header", shortUDPLen, 20 + 6},
|
|
||||||
}
|
|
||||||
for _, tt := range untouched {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
if tt.cksum >= 0 {
|
|
||||||
binary.BigEndian.PutUint16(tt.pkt[tt.cksum:], 0x1234)
|
|
||||||
}
|
|
||||||
want := append([]byte(nil), tt.pkt...)
|
|
||||||
SetTransportChecksum(tt.pkt)
|
|
||||||
assert.Equal(t, want, tt.pkt)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Run("too short to carry a header", func(t *testing.T) {
|
|
||||||
for _, pkt := range [][]byte{nil, {}, {0x45}, {0x60}} {
|
|
||||||
assert.NotPanics(t, func() { SetTransportChecksum(pkt) })
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("tcp checksum of zero goes out as zero", func(t *testing.T) {
|
|
||||||
pkt := serialize(t, ip4tcp, tcp(ip4tcp), gopacket.Payload{0, 0})
|
|
||||||
c := binary.BigEndian.Uint16(pkt[36:38])
|
|
||||||
require.NotZero(t, c)
|
|
||||||
// Only udp reserves zero to mean "not computed", so tcp keeps it.
|
|
||||||
binary.BigEndian.PutUint16(pkt[40:42], c)
|
|
||||||
SetTransportChecksum(pkt)
|
|
||||||
assert.Zero(t, binary.BigEndian.Uint16(pkt[36:38]))
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("udp checksum of zero goes out as 0xffff", func(t *testing.T) {
|
|
||||||
pkt := serialize(t, ip4udp, udp(ip4udp), gopacket.Payload{0, 0})
|
|
||||||
c := binary.BigEndian.Uint16(pkt[26:28])
|
|
||||||
require.NotZero(t, c)
|
|
||||||
// The one's complement sum is now 0xffff - c; adding c to the payload
|
|
||||||
// makes it 0xffff, whose complement is zero.
|
|
||||||
binary.BigEndian.PutUint16(pkt[28:30], c)
|
|
||||||
SetTransportChecksum(pkt)
|
|
||||||
assert.Equal(t, uint16(0xffff), binary.BigEndian.Uint16(pkt[26:28]))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFold(t *testing.T) {
|
|
||||||
// 0xffff is the fold's fixed point, so a loop bound one notch tight never
|
|
||||||
// terminates on it.
|
|
||||||
for _, tt := range []struct {
|
|
||||||
in uint32
|
|
||||||
want uint16
|
|
||||||
}{
|
|
||||||
{0, 0},
|
|
||||||
{0xffff, 0xffff},
|
|
||||||
{0x10000, 1},
|
|
||||||
{0x1fffe, 0xffff},
|
|
||||||
{0xffffffff, 0xffff},
|
|
||||||
} {
|
|
||||||
assert.Equal(t, tt.want, fold(tt.in))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+9
-41
@@ -2,16 +2,11 @@ package iputil
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"errors"
|
|
||||||
|
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
"golang.org/x/net/ipv6"
|
"golang.org/x/net/ipv6"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ErrIPv6CouldNotFindPayload is returned when the ipv6 extension header chain is truncated before a terminal
|
|
||||||
// upper layer protocol is reached.
|
|
||||||
var ErrIPv6CouldNotFindPayload = errors.New("could not find payload in ipv6 packet")
|
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// MaxIPv4RejectPacketSize is the largest IPv4 reject packet:
|
// MaxIPv4RejectPacketSize is the largest IPv4 reject packet:
|
||||||
// - 20 byte ipv4 header
|
// - 20 byte ipv4 header
|
||||||
@@ -27,13 +22,6 @@ const (
|
|||||||
maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000
|
maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000
|
||||||
|
|
||||||
MaxRejectPacketSize = maxIPv6RejectPacketSize
|
MaxRejectPacketSize = maxIPv6RejectPacketSize
|
||||||
|
|
||||||
IPProtocolICMP = 1
|
|
||||||
IPProtocolICMPv6 = 58
|
|
||||||
IPProtocolTCP = 6
|
|
||||||
IPProtocolUDP = 17
|
|
||||||
ICMPv6TypeEchoRequest = 128
|
|
||||||
ICMPv6TypeEchoReply = 129
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func CreateRejectPacket(packet []byte, out []byte) []byte {
|
func CreateRejectPacket(packet []byte, out []byte) []byte {
|
||||||
@@ -211,8 +199,8 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
|
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
|
||||||
proto, offset, isFragment, _, err := IPv6FindUpperProtocol(packet)
|
proto, offset, isFragment := ipv6FindUpperProtocol(packet)
|
||||||
if err != nil || isFragment {
|
if isFragment {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
switch proto {
|
switch proto {
|
||||||
@@ -345,60 +333,40 @@ func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte {
|
|||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
// IPv6FindUpperProtocol walks the ipv6 extension header chain and returns the upper layer protocol, the
|
func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
|
||||||
// offset it begins at, and whether the packet is a non-first fragment. Only the RFC 8200 and IANA extension
|
|
||||||
// headers below are walked. Everything else, including Mobility (135), HIP (139), Shim6 (140), experimental
|
|
||||||
// 253/254, and real upper layer protocols like SCTP or GRE, is terminal. Walking those as extension headers
|
|
||||||
// is a firewall bypass, so they fail closed. For a non-first fragment the returned protocol is the fragmented
|
|
||||||
// protocol and offset points at the fragment header, there is no transport header to locate. Returns
|
|
||||||
// ErrIPv6CouldNotFindPayload if packet is smaller than an ipv6 header or the chain is truncated before a
|
|
||||||
// terminal protocol is reached.
|
|
||||||
func IPv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool, anyFragment bool, err error) {
|
|
||||||
const maxIPv6ExtHeaders = 8
|
|
||||||
if len(packet) < ipv6.HeaderLen {
|
|
||||||
return 0, 0, false, false, ErrIPv6CouldNotFindPayload
|
|
||||||
}
|
|
||||||
nextHeader = packet[6]
|
nextHeader = packet[6]
|
||||||
offset = ipv6.HeaderLen
|
offset = ipv6.HeaderLen
|
||||||
|
|
||||||
for range maxIPv6ExtHeaders {
|
for {
|
||||||
switch nextHeader {
|
switch nextHeader {
|
||||||
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
|
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
|
||||||
if len(packet) < offset+2 {
|
if len(packet) < offset+2 {
|
||||||
return nextHeader, offset, isFragment, anyFragment, ErrIPv6CouldNotFindPayload
|
return nextHeader, offset, isFragment
|
||||||
}
|
}
|
||||||
nextHeader = packet[offset]
|
nextHeader = packet[offset]
|
||||||
offset += (int(packet[offset+1]) + 1) << 3
|
offset += (int(packet[offset+1]) + 1) << 3
|
||||||
|
|
||||||
case 44: // Fragment
|
case 44: // Fragment
|
||||||
if len(packet) < offset+8 {
|
if len(packet) < offset+8 {
|
||||||
return nextHeader, offset, isFragment, anyFragment, ErrIPv6CouldNotFindPayload
|
return nextHeader, offset, isFragment
|
||||||
}
|
}
|
||||||
anyFragment = true
|
|
||||||
// Non-first fragments carry no transport header, report the fragmented protocol and stop
|
|
||||||
if packet[offset+2] != 0 || packet[offset+3]&0xf8 != 0 {
|
if packet[offset+2] != 0 || packet[offset+3]&0xf8 != 0 {
|
||||||
return packet[offset], offset, true, anyFragment, nil
|
isFragment = true
|
||||||
}
|
}
|
||||||
nextHeader = packet[offset]
|
nextHeader = packet[offset]
|
||||||
offset += 8
|
offset += 8
|
||||||
|
|
||||||
case 51: // AH
|
case 51: // AH
|
||||||
if len(packet) < offset+2 {
|
if len(packet) < offset+2 {
|
||||||
return nextHeader, offset, isFragment, anyFragment, ErrIPv6CouldNotFindPayload
|
return nextHeader, offset, isFragment
|
||||||
}
|
}
|
||||||
nextHeader = packet[offset]
|
nextHeader = packet[offset]
|
||||||
offset += (int(packet[offset+1]) + 2) << 2
|
offset += (int(packet[offset+1]) + 2) << 2
|
||||||
|
|
||||||
default:
|
default:
|
||||||
// A prior extension header can declare a length that advances offset past the packet. The terminal
|
return nextHeader, offset, isFragment
|
||||||
// protocol's header isn't actually here, so treat the chain as truncated rather than classifying it.
|
|
||||||
if offset > len(packet) {
|
|
||||||
return nextHeader, offset, isFragment, anyFragment, ErrIPv6CouldNotFindPayload
|
|
||||||
}
|
|
||||||
return nextHeader, offset, isFragment, anyFragment, nil
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nextHeader, offset, isFragment, anyFragment, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
||||||
|
|||||||
@@ -1,13 +1,11 @@
|
|||||||
package iputil
|
package iputil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"net"
|
"net"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
"golang.org/x/net/ipv6"
|
"golang.org/x/net/ipv6"
|
||||||
)
|
)
|
||||||
@@ -181,46 +179,6 @@ func Test_CreateRejectPacket_NoICMPError(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Test_CreateRejectPacket_RespectsCap ensures it is impossible for
|
|
||||||
// an oversized ICMPv6 reject to overwrite the neighbor segment's bytes.
|
|
||||||
func Test_CreateRejectPacket_RespectsCap(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
|
|
||||||
// Inner IPv6 UDP packet. An ICMPv6 reject copies the whole inner packet
|
|
||||||
// plus a 48-byte header (40 IPv6 + 8 ICMPv6), so it needs 48 more bytes
|
|
||||||
// than the inner packet length.
|
|
||||||
inner := makeIPv6Packet(src, dst, 17, make([]byte, 20))
|
|
||||||
|
|
||||||
// The ciphertext scratch reused as the reject buffer is the received
|
|
||||||
// datagram: 16-byte Nebula header + inner + 16-byte AEAD tag. That is only
|
|
||||||
// 32 bytes of slack, so a full ICMPv6 reject overruns it by 16 bytes.
|
|
||||||
const nebulaOverhead = 32
|
|
||||||
segLen := len(inner) + nebulaOverhead
|
|
||||||
|
|
||||||
// Shared backing row laid out as [segment][neighbor's 16-byte Nebula header].
|
|
||||||
const neighborHdr = 16
|
|
||||||
sentinel := bytes.Repeat([]byte{0xAB}, neighborHdr)
|
|
||||||
|
|
||||||
// Uncapped: the slice's capacity reaches into the neighbor, reproducing
|
|
||||||
// the overrun that silently drops the neighbor packet.
|
|
||||||
backing := make([]byte, segLen+neighborHdr)
|
|
||||||
copy(backing[segLen:], sentinel)
|
|
||||||
reject := CreateRejectPacket(inner, backing[:segLen])
|
|
||||||
assert.NotNil(t, reject, "uncapped buffer reaches into the neighbor, so the reject is built")
|
|
||||||
assert.NotEqual(t, sentinel, backing[segLen:segLen+neighborHdr],
|
|
||||||
"without the cap the oversized reject overruns into the neighbor segment")
|
|
||||||
|
|
||||||
// Capped (the fix): cap==len, so the builder cannot exceed the segment. The
|
|
||||||
// reject does not fit, so it is refused rather than corrupting the neighbor.
|
|
||||||
backing = make([]byte, segLen+neighborHdr)
|
|
||||||
copy(backing[segLen:], sentinel)
|
|
||||||
reject = CreateRejectPacket(inner, backing[:segLen:segLen])
|
|
||||||
assert.Nil(t, reject, "capped segment is 16 bytes too small for a full ICMPv6 reject, so it is refused")
|
|
||||||
assert.Equal(t, sentinel, backing[segLen:segLen+neighborHdr],
|
|
||||||
"capped segment must leave the neighbor untouched")
|
|
||||||
}
|
|
||||||
|
|
||||||
func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte {
|
func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte {
|
||||||
b := make([]byte, ipv6.HeaderLen+len(payload))
|
b := make([]byte, ipv6.HeaderLen+len(payload))
|
||||||
b[0] = ipv6.Version << 4
|
b[0] = ipv6.Version << 4
|
||||||
@@ -516,63 +474,3 @@ func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) {
|
|||||||
result := CreateICMPEchoResponse(packet, out)
|
result := CreateICMPEchoResponse(packet, out)
|
||||||
assert.Nil(t, result)
|
assert.Nil(t, result)
|
||||||
}
|
}
|
||||||
|
|
||||||
func Test_IPv6FindUpperProtocol(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
|
|
||||||
// 8 byte extension/transport stand-ins, first byte is the next header, second is the length field
|
|
||||||
extToTCP := []byte{6, 0, 0, 0, 0, 0, 0, 0} // len 0 -> 8 bytes, next = TCP
|
|
||||||
extToUDP := []byte{17, 0, 0, 0, 0, 0, 0, 0} // len 0 -> 8 bytes, next = UDP
|
|
||||||
extToRouting := []byte{43, 0, 0, 0, 0, 0, 0, 0} // len 0 -> 8 bytes, next = Routing
|
|
||||||
ahToUDP := []byte{17, 0, 0, 0, 0, 0, 0, 0} // AH len 0 -> (0+2)<<2 = 8 bytes, next = UDP
|
|
||||||
firstFragToUDP := []byte{17, 0, 0, 1, 0, 0, 0, 1} // frag offset 0, M=1, next = UDP
|
|
||||||
nonFirstFrag := []byte{17, 0, 0, 9, 0, 0, 0, 1} // frag offset non-zero, next = UDP
|
|
||||||
transport := []byte{0, 80, 1, 187, 0, 0, 0, 0} // stand-in bytes, IPv6FindUpperProtocol never reads ports
|
|
||||||
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
nextHeader uint8
|
|
||||||
payload []byte
|
|
||||||
wantProto uint8
|
|
||||||
wantOffset int
|
|
||||||
wantFragment bool
|
|
||||||
wantAnyFrag bool
|
|
||||||
wantErr error
|
|
||||||
}{
|
|
||||||
{"plain udp", 17, transport, 17, ipv6.HeaderLen, false, false, nil},
|
|
||||||
{"hop-by-hop then tcp", 0, append(extToTCP, transport...), 6, ipv6.HeaderLen + 8, false, false, nil},
|
|
||||||
{"routing then tcp", 43, append(extToTCP, transport...), 6, ipv6.HeaderLen + 8, false, false, nil},
|
|
||||||
{"destination then udp", 60, append(extToUDP, transport...), 17, ipv6.HeaderLen + 8, false, false, nil},
|
|
||||||
{"hop-by-hop, routing, then tcp", 0, append(append(extToRouting, extToTCP...), transport...), 6, ipv6.HeaderLen + 16, false, false, nil},
|
|
||||||
{"ah then udp", 51, append(ahToUDP, transport...), 17, ipv6.HeaderLen + 8, false, false, nil},
|
|
||||||
{"first fragment walks to transport", 44, append(firstFragToUDP, transport...), 17, ipv6.HeaderLen + 8, false, true, nil},
|
|
||||||
{"non-first fragment stops", 44, append(nonFirstFrag, transport...), 17, ipv6.HeaderLen, true, true, nil},
|
|
||||||
{"unknown protocol is terminal", 132, transport, 132, ipv6.HeaderLen, false, false, nil}, // SCTP
|
|
||||||
{"truncated extension header", 0, nil, 0, ipv6.HeaderLen, false, false, ErrIPv6CouldNotFindPayload},
|
|
||||||
// Destination Options with a declared length (255+1)*8 = 2048 that runs past the 48 byte buffer, next = SCTP
|
|
||||||
{"extension length past buffer", 60, []byte{132, 255, 0, 0, 0, 0, 0, 0}, 132, ipv6.HeaderLen + 2048, false, false, ErrIPv6CouldNotFindPayload},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
packet := makeIPv6Packet(src, dst, tt.nextHeader, tt.payload)
|
|
||||||
proto, offset, isFragment, anyFragment, err := IPv6FindUpperProtocol(packet)
|
|
||||||
if tt.wantErr != nil {
|
|
||||||
assert.ErrorIs(t, err, tt.wantErr)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, tt.wantProto, proto)
|
|
||||||
assert.Equal(t, tt.wantOffset, offset)
|
|
||||||
assert.Equal(t, tt.wantFragment, isFragment)
|
|
||||||
assert.Equal(t, tt.wantAnyFrag, anyFragment)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// A packet smaller than an ipv6 header must error rather than panic reading byte 6
|
|
||||||
t.Run("shorter than ipv6 header", func(t *testing.T) {
|
|
||||||
_, _, _, _, err := IPv6FindUpperProtocol(make([]byte, 6))
|
|
||||||
assert.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|||||||
+1
-15
@@ -34,9 +34,7 @@ type LightHouse struct {
|
|||||||
|
|
||||||
myVpnNetworks []netip.Prefix
|
myVpnNetworks []netip.Prefix
|
||||||
myVpnNetworksTable *bart.Lite
|
myVpnNetworksTable *bart.Lite
|
||||||
// myVpnAddrsTable contains our overlay host addrs, as opposed to the overlay networks
|
punchy *Punchy
|
||||||
myVpnAddrsTable *bart.Lite
|
|
||||||
punchy *Punchy
|
|
||||||
|
|
||||||
// localAddrsFn enumerates the underlay addresses we advertise. It is a field so tests can supply simulated
|
// 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.
|
// addresses rather than whatever this machine's NICs happen to be. Set it before Start.
|
||||||
@@ -106,7 +104,6 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
|||||||
amLighthouse: amLighthouse,
|
amLighthouse: amLighthouse,
|
||||||
myVpnNetworks: cs.myVpnNetworks,
|
myVpnNetworks: cs.myVpnNetworks,
|
||||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||||
myVpnAddrsTable: cs.myVpnAddrsTable,
|
|
||||||
addrMap: make(map[netip.Addr]*RemoteList),
|
addrMap: make(map[netip.Addr]*RemoteList),
|
||||||
nebulaPort: nebulaPort,
|
nebulaPort: nebulaPort,
|
||||||
punchy: p,
|
punchy: p,
|
||||||
@@ -1161,17 +1158,6 @@ func (lhh *LightHouseHandler) handleHostQuery(n *NebulaMeta, fromVpnAddrs []neti
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Don't respond to requests for us.
|
|
||||||
if lhh.lh.myVpnAddrsTable.Contains(queryVpnAddr) {
|
|
||||||
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
lhh.l.Debug("Ignoring HostQuery for one of my own addresses",
|
|
||||||
"fromVpnAddrs", fromVpnAddrs,
|
|
||||||
"queryVpnAddr", queryVpnAddr,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
found, ln, err := lhh.lh.queryAndPrepMessage(queryVpnAddr, func(c *cache) (int, error) {
|
found, ln, err := lhh.lh.queryAndPrepMessage(queryVpnAddr, func(c *cache) (int, error) {
|
||||||
n = lhh.resetMeta()
|
n = lhh.resetMeta()
|
||||||
n.Type = NebulaMeta_HostQueryReply
|
n.Type = NebulaMeta_HostQueryReply
|
||||||
|
|||||||
+54
-80
@@ -27,27 +27,15 @@ func TestOldIPv4Only(t *testing.T) {
|
|||||||
assert.Equal(t, binary.BigEndian.Uint32(bp[:]), m.GetAddr())
|
assert.Equal(t, binary.BigEndian.Uint32(bp[:]), m.GetAddr())
|
||||||
}
|
}
|
||||||
|
|
||||||
func testCertState(networks ...netip.Prefix) *CertState {
|
|
||||||
cs := &CertState{
|
|
||||||
myVpnNetworks: networks,
|
|
||||||
myVpnNetworksTable: new(bart.Lite),
|
|
||||||
myVpnAddrs: make([]netip.Addr, 0, len(networks)),
|
|
||||||
myVpnAddrsTable: new(bart.Lite),
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, n := range networks {
|
|
||||||
cs.myVpnNetworksTable.Insert(n)
|
|
||||||
cs.myVpnAddrs = append(cs.myVpnAddrs, n.Addr())
|
|
||||||
cs.myVpnAddrsTable.Insert(netip.PrefixFrom(n.Addr(), n.Addr().BitLen()))
|
|
||||||
}
|
|
||||||
|
|
||||||
return cs
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_lhStaticMapping(t *testing.T) {
|
func Test_lhStaticMapping(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
||||||
cs := testCertState(myVpnNet)
|
nt := new(bart.Lite)
|
||||||
|
nt.Insert(myVpnNet)
|
||||||
|
cs := &CertState{
|
||||||
|
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||||
|
myVpnNetworksTable: nt,
|
||||||
|
}
|
||||||
lh1 := "10.128.0.2"
|
lh1 := "10.128.0.2"
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
@@ -67,7 +55,12 @@ func Test_lhStaticMapping(t *testing.T) {
|
|||||||
func TestReloadLighthouseInterval(t *testing.T) {
|
func TestReloadLighthouseInterval(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
||||||
cs := testCertState(myVpnNet)
|
nt := new(bart.Lite)
|
||||||
|
nt.Insert(myVpnNet)
|
||||||
|
cs := &CertState{
|
||||||
|
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||||
|
myVpnNetworksTable: nt,
|
||||||
|
}
|
||||||
lh1 := "10.128.0.2"
|
lh1 := "10.128.0.2"
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
@@ -97,7 +90,12 @@ func TestReloadLighthouseInterval(t *testing.T) {
|
|||||||
func BenchmarkLighthouseHandleRequest(b *testing.B) {
|
func BenchmarkLighthouseHandleRequest(b *testing.B) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/0")
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/0")
|
||||||
cs := testCertState(myVpnNet)
|
nt := new(bart.Lite)
|
||||||
|
nt.Insert(myVpnNet)
|
||||||
|
cs := &CertState{
|
||||||
|
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||||
|
myVpnNetworksTable: nt,
|
||||||
|
}
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
lh, err := NewLightHouseFromConfig(b.Context(), l, c, cs, nil, nil)
|
lh, err := NewLightHouseFromConfig(b.Context(), l, c, cs, nil, nil)
|
||||||
@@ -197,7 +195,12 @@ func TestLighthouse_Memory(t *testing.T) {
|
|||||||
c.Settings["listen"] = map[string]any{"port": 4242}
|
c.Settings["listen"] = map[string]any{"port": 4242}
|
||||||
|
|
||||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||||
cs := testCertState(myVpnNet)
|
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)
|
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||||
lh.ifce = &mockEncWriter{}
|
lh.ifce = &mockEncWriter{}
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -277,7 +280,12 @@ func TestLighthouse_reload(t *testing.T) {
|
|||||||
c.Settings["listen"] = map[string]any{"port": 4242}
|
c.Settings["listen"] = map[string]any{"port": 4242}
|
||||||
|
|
||||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||||
cs := testCertState(myVpnNet)
|
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)
|
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -307,7 +315,12 @@ func TestLighthouse_reloadStaticHostMap(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||||
cs := testCertState(myVpnNet)
|
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)
|
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -416,9 +429,7 @@ func TestLighthouse_reloadStaticHostMap(t *testing.T) {
|
|||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
||||||
}
|
}
|
||||||
|
|
||||||
// sendLHHostRequest delivers a HostQuery to lhh and hands back the writer that
|
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
||||||
// captured what it emitted. Pass a nil filter to see every message.
|
|
||||||
func sendLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler, filter *NebulaMeta_MessageType) *testEncWriter {
|
|
||||||
req := &NebulaMeta{
|
req := &NebulaMeta{
|
||||||
Type: NebulaMeta_HostQuery,
|
Type: NebulaMeta_HostQuery,
|
||||||
Details: &NebulaMetaDetails{},
|
Details: &NebulaMetaDetails{},
|
||||||
@@ -436,59 +447,12 @@ func sendLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr,
|
|||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
w := &testEncWriter{metaFilter: filter}
|
|
||||||
lhh.HandleRequest(fromAddr, []netip.Addr{myVpnIp}, b, w)
|
|
||||||
return w
|
|
||||||
}
|
|
||||||
|
|
||||||
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
|
||||||
filter := NebulaMeta_HostQueryReply
|
filter := NebulaMeta_HostQueryReply
|
||||||
return sendLHHostRequest(fromAddr, myVpnIp, queryVpnIp, lhh, &filter).lastReply
|
w := &testEncWriter{
|
||||||
}
|
metaFilter: &filter,
|
||||||
|
|
||||||
func TestLighthouse_IgnoresHostQueryForItself(t *testing.T) {
|
|
||||||
// Validate that we don't answer host queries for our own address.
|
|
||||||
l := test.NewLogger()
|
|
||||||
|
|
||||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
|
||||||
myVpnIp := myVpnNet.Addr()
|
|
||||||
|
|
||||||
c := config.NewC(l)
|
|
||||||
c.Settings["lighthouse"] = map[string]any{"am_lighthouse": true}
|
|
||||||
c.Settings["listen"] = map[string]any{"port": 4242}
|
|
||||||
// Add a static_host_map entry for ourselves, so our address
|
|
||||||
// is in the addrMap.
|
|
||||||
c.Settings["static_host_map"] = map[string]any{
|
|
||||||
myVpnIp.String(): []any{"192.168.100.1:4242"},
|
|
||||||
}
|
}
|
||||||
|
lhh.HandleRequest(fromAddr, []netip.Addr{myVpnIp}, b, w)
|
||||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, testCertState(myVpnNet), nil, nil)
|
return w.lastReply
|
||||||
require.NoError(t, err)
|
|
||||||
lh.ifce = &mockEncWriter{}
|
|
||||||
lhh := lh.NewRequestHandler()
|
|
||||||
|
|
||||||
peerVpnIp := netip.MustParseAddr("10.128.0.2")
|
|
||||||
peerUdpAddr := netip.MustParseAddrPort("10.0.0.2:4242")
|
|
||||||
otherVpnIp := netip.MustParseAddr("10.128.0.3")
|
|
||||||
otherUdpAddr := netip.MustParseAddrPort("10.0.0.3:4242")
|
|
||||||
|
|
||||||
newLHHostUpdate(peerUdpAddr, peerVpnIp, []netip.AddrPort{peerUdpAddr}, lhh)
|
|
||||||
newLHHostUpdate(otherUdpAddr, otherVpnIp, []netip.AddrPort{otherUdpAddr}, lhh)
|
|
||||||
|
|
||||||
// Control: a query about a real peer is still answered, and still ends with
|
|
||||||
// the punch notification aimed at the host that was asked about.
|
|
||||||
w := sendLHHostRequest(peerUdpAddr, peerVpnIp, otherVpnIp, lhh, nil)
|
|
||||||
require.NotNil(t, w.lastReply.msg)
|
|
||||||
assert.Equal(t, NebulaMeta_HostPunchNotification, w.lastReply.msg.Type)
|
|
||||||
assert.Equal(t, otherVpnIp, w.lastReply.vpnIp)
|
|
||||||
|
|
||||||
// Now validate that we don't send to ourselves.
|
|
||||||
found, _, err := lh.queryAndPrepMessage(myVpnIp, func(*cache) (int, error) { return 0, nil })
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.True(t, found, "the lighthouse should hold a cache entry for its own address")
|
|
||||||
|
|
||||||
w = sendLHHostRequest(peerUdpAddr, peerVpnIp, myVpnIp, lhh, nil)
|
|
||||||
assert.Nil(t, w.lastReply.msg, "a query about our own address must produce no reply and no punch notification")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func newLHHostUpdate(fromAddr netip.AddrPort, vpnIp netip.Addr, addrs []netip.AddrPort, lhh *LightHouseHandler) {
|
func newLHHostUpdate(fromAddr netip.AddrPort, vpnIp netip.Addr, addrs []netip.AddrPort, lhh *LightHouseHandler) {
|
||||||
@@ -534,7 +498,7 @@ type testEncWriter struct {
|
|||||||
protocolVersion cert.Version
|
protocolVersion cert.Version
|
||||||
}
|
}
|
||||||
|
|
||||||
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool) {
|
||||||
}
|
}
|
||||||
func (tw *testEncWriter) Handshake(vpnIp netip.Addr) {
|
func (tw *testEncWriter) Handshake(vpnIp netip.Addr) {
|
||||||
}
|
}
|
||||||
@@ -678,7 +642,12 @@ func TestLighthouse_Dont_Delete_Static_Hosts(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||||
cs := testCertState(myVpnNet)
|
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)
|
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
lh.ifce = &mockEncWriter{}
|
lh.ifce = &mockEncWriter{}
|
||||||
@@ -739,7 +708,12 @@ func TestLighthouse_DeletesWork(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||||
cs := testCertState(myVpnNet)
|
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)
|
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
lh.ifce = &mockEncWriter{}
|
lh.ifce = &mockEncWriter{}
|
||||||
|
|||||||
@@ -6,15 +6,11 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
|
||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
"slices"
|
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/cpupick"
|
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
"github.com/slackhq/nebula/sshd"
|
"github.com/slackhq/nebula/sshd"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
@@ -24,12 +20,6 @@ import (
|
|||||||
|
|
||||||
type m = map[string]any
|
type m = map[string]any
|
||||||
|
|
||||||
// maxRoutines caps routines below the RejectHeadroom nonce gap so concurrent senders can't race the counter past wrap.
|
|
||||||
const maxRoutines = 1 << 16
|
|
||||||
|
|
||||||
// The reject headroom must exceed every sender that can be mid-reservation at once, about two per routine.
|
|
||||||
const _ = noiseutil.RejectHeadroom - 4*maxRoutines
|
|
||||||
|
|
||||||
func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, deviceFactory overlay.DeviceFactory) (retcon *Control, reterr error) {
|
func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, deviceFactory overlay.DeviceFactory) (retcon *Control, reterr error) {
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
// Automatically cancel the context if Main returns an error, to signal all created goroutines to quit.
|
// Automatically cancel the context if Main returns an error, to signal all created goroutines to quit.
|
||||||
@@ -43,9 +33,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
buildVersion = moduleVersion()
|
buildVersion = moduleVersion()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Debug builds (-tags debug) serve pprof on :6060; a no-op otherwise.
|
|
||||||
startPprofServer(ctx, l)
|
|
||||||
|
|
||||||
// Print the config if in test, the exit comes later
|
// Print the config if in test, the exit comes later
|
||||||
if configTest {
|
if configTest {
|
||||||
b, err := yaml.Marshal(c.Settings)
|
b, err := yaml.Marshal(c.Settings)
|
||||||
@@ -94,6 +81,9 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
if routines < 1 {
|
if routines < 1 {
|
||||||
routines = 1
|
routines = 1
|
||||||
}
|
}
|
||||||
|
if routines > 1 {
|
||||||
|
l.Info("Using multiple routines", "routines", routines)
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
// deprecated and undocumented
|
// deprecated and undocumented
|
||||||
tunQueues := c.GetInt("tun.routines", 1)
|
tunQueues := c.GetInt("tun.routines", 1)
|
||||||
@@ -103,12 +93,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
l.Warn("Setting tun.routines and listen.routines is deprecated. Use `routines` instead", "routines", routines)
|
l.Warn("Setting tun.routines and listen.routines is deprecated. Use `routines` instead", "routines", routines)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if routines > maxRoutines {
|
|
||||||
l.Warn("Using multiple routines", "routines", maxRoutines, "clamped", true, "requestedRoutines", routines)
|
|
||||||
routines = maxRoutines
|
|
||||||
} else if routines > 1 {
|
|
||||||
l.Info("Using multiple routines", "routines", routines)
|
|
||||||
}
|
|
||||||
|
|
||||||
// EXPERIMENTAL
|
// EXPERIMENTAL
|
||||||
// Intentionally not documented yet while we do more testing and determine
|
// Intentionally not documented yet while we do more testing and determine
|
||||||
@@ -176,21 +160,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
}
|
}
|
||||||
|
|
||||||
for i := 0; i < routines; i++ {
|
for i := 0; i < routines; i++ {
|
||||||
listen := netip.AddrPortFrom(listenHost, uint16(port))
|
l.Info("listening", "addr", netip.AddrPortFrom(listenHost, uint16(port)))
|
||||||
l.Info("listening", "addr", listen)
|
udpServer, err := udp.NewListener(l, listenHost, port, routines > 1, c.GetInt("listen.batch", 64))
|
||||||
batchSize := c.GetInt("listen.batch", 64)
|
|
||||||
if batchSize < 1 {
|
|
||||||
oldBatch := batchSize
|
|
||||||
batchSize = 1
|
|
||||||
l.Warn("listen.batch size is invalid", "provided", oldBatch, "overridden to", batchSize)
|
|
||||||
}
|
|
||||||
udpSettings := udp.Settings{
|
|
||||||
Listen: listen,
|
|
||||||
Multi: routines > 1,
|
|
||||||
Batch: batchSize,
|
|
||||||
Offloads: c.GetBool("listen.udp_offloads", false),
|
|
||||||
}
|
|
||||||
udpServer, err := udp.NewListener(l, udpSettings)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, util.NewContextualError("Failed to open udp listener", m{"queue": i}, err)
|
return nil, util.NewContextualError("Failed to open udp listener", m{"queue": i}, err)
|
||||||
}
|
}
|
||||||
@@ -239,37 +210,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
l.Warn("Failed to start DNS responder", "error", err)
|
l.Warn("Failed to start DNS responder", "error", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
pinThreads := c.GetBool("tun.pin_threads", true)
|
|
||||||
cpuAffinity := parseCpuAffinity(c, l, routines)
|
|
||||||
if pinThreads && routines > 1 && len(cpuAffinity) == 0 && !configTest {
|
|
||||||
// The operator didn't choose pin CPUs, so pick a default set that
|
|
||||||
// prefers performance cores and doesn't stack co-located instances
|
|
||||||
// onto allowed[0].
|
|
||||||
|
|
||||||
// key is used to seed the spreading of routines->cores.
|
|
||||||
// use PID if you want to ensure many different Nebulas in VMs or containers land on different cores
|
|
||||||
// use port if you want to always end up on the same cores, ideal for benchmarking.
|
|
||||||
key := uint64(os.Getpid()) //default to PID
|
|
||||||
pinKeyStr := strings.ToLower(c.GetString("tun.pin_threads_key", ""))
|
|
||||||
switch pinKeyStr {
|
|
||||||
case "":
|
|
||||||
l.Debug("tun.pin_threads_key is empty, using PID")
|
|
||||||
case "pid":
|
|
||||||
l.Debug("tun.pin_threads_key is PID")
|
|
||||||
case "port":
|
|
||||||
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
|
|
||||||
l.Info("tun.pin_threads_key is port number")
|
|
||||||
key = uint64(ap.Port())
|
|
||||||
} else {
|
|
||||||
l.Warn("Failed to get a port number for tun.pin_threads_key, falling back to PID", "err", err)
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
l.Warn("tun.pin_threads_key is invalid, using PID")
|
|
||||||
}
|
|
||||||
|
|
||||||
cpuAffinity = cpupick.Default(routines, key, l)
|
|
||||||
}
|
|
||||||
|
|
||||||
ifConfig := &InterfaceConfig{
|
ifConfig := &InterfaceConfig{
|
||||||
HostMap: hostMap,
|
HostMap: hostMap,
|
||||||
Inside: tun,
|
Inside: tun,
|
||||||
@@ -291,8 +231,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
||||||
punchy: punchy,
|
punchy: punchy,
|
||||||
ConntrackCacheTimeout: conntrackCacheTimeout,
|
ConntrackCacheTimeout: conntrackCacheTimeout,
|
||||||
CpuAffinity: cpuAffinity,
|
|
||||||
PinThreads: pinThreads,
|
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -347,70 +285,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// parseCpuAffinity reads `tun.cpu_affinity` from the config — a list of
|
|
||||||
// integer CPU IDs, one per TUN reader goroutine. Empty / unset returns nil
|
|
||||||
// (listenIn falls back to spreading queues across the allowed CPU set).
|
|
||||||
// Length mismatch with `routines` is a warning, not an error: shorter lists
|
|
||||||
// are modulo-cycled across queues, longer lists' tail is ignored. Invalid
|
|
||||||
// entries (non-integer, or a CPU ID we're not allowed to run on) are also a
|
|
||||||
// warning and disable the override entirely so we don't silently pin to the
|
|
||||||
// wrong CPU. Entries are validated against the process's current affinity
|
|
||||||
// mask (util.AllowedCPUs) rather than 0..NumCPU-1: under a cgroup cpuset or
|
|
||||||
// taskset the runnable IDs are frequently not that contiguous range, and
|
|
||||||
// pinning to an unrunnable ID always fails. If the allowed set can't be
|
|
||||||
// determined we fall back to a plain non-negative check.
|
|
||||||
func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
|
|
||||||
raw := c.Get("tun.cpu_affinity")
|
|
||||||
if raw == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
rv, ok := raw.([]any)
|
|
||||||
if !ok {
|
|
||||||
l.Warn("tun.cpu_affinity must be a list of integers; ignoring", "value", raw)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
// allowed is the set of CPU IDs we're actually permitted to run on. A nil
|
|
||||||
// slice (unsupported platform or lookup error) means "can't tell", so we
|
|
||||||
// only apply the weaker non-negative check in that case.
|
|
||||||
allowed, err := util.AllowedCPUs()
|
|
||||||
if err != nil {
|
|
||||||
l.Warn("could not determine allowed CPUs; validating tun.cpu_affinity against non-negative only", "error", err)
|
|
||||||
allowed = nil
|
|
||||||
}
|
|
||||||
cpus := make([]int, 0, len(rv))
|
|
||||||
for i, e := range rv {
|
|
||||||
var cpu int
|
|
||||||
switch v := e.(type) {
|
|
||||||
case int:
|
|
||||||
cpu = v
|
|
||||||
case int64:
|
|
||||||
cpu = int(v)
|
|
||||||
case float64:
|
|
||||||
cpu = int(v)
|
|
||||||
default:
|
|
||||||
l.Warn("tun.cpu_affinity entry not an integer; ignoring affinity",
|
|
||||||
"index", i, "value", e)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
if cpu < 0 {
|
|
||||||
l.Warn("tun.cpu_affinity entry out of range; ignoring affinity",
|
|
||||||
"index", i, "cpu", cpu)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
if len(allowed) > 0 && !slices.Contains(allowed, cpu) {
|
|
||||||
l.Warn("tun.cpu_affinity entry not in allowed CPU set; ignoring affinity",
|
|
||||||
"index", i, "cpu", cpu, "allowed", allowed)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
cpus = append(cpus, cpu)
|
|
||||||
}
|
|
||||||
if len(cpus) != routines {
|
|
||||||
l.Warn("tun.cpu_affinity length doesn't match routines; queues will modulo-cycle through the list",
|
|
||||||
"affinity_len", len(cpus), "routines", routines)
|
|
||||||
}
|
|
||||||
return cpus
|
|
||||||
}
|
|
||||||
|
|
||||||
func moduleVersion() string {
|
func moduleVersion() string {
|
||||||
info, ok := debug.ReadBuildInfo()
|
info, ok := debug.ReadBuildInfo()
|
||||||
if !ok {
|
if !ok {
|
||||||
|
|||||||
@@ -1,51 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"github.com/slackhq/nebula/util"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestParseCpuAffinity(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
|
|
||||||
// newConfig returns a config.C with tun.cpu_affinity set to v. A nil v
|
|
||||||
// leaves the key unset.
|
|
||||||
newConfig := func(v any) *config.C {
|
|
||||||
c := config.NewC(l)
|
|
||||||
if v != nil {
|
|
||||||
c.Settings["tun"] = map[string]any{"cpu_affinity": v}
|
|
||||||
}
|
|
||||||
return c
|
|
||||||
}
|
|
||||||
|
|
||||||
// unset -> nil (listenIn falls back to spreading across the allowed set)
|
|
||||||
assert.Nil(t, parseCpuAffinity(newConfig(nil), l, 1))
|
|
||||||
|
|
||||||
// Pick a CPU we're actually allowed to run on so a valid list survives
|
|
||||||
// validation regardless of the host's affinity mask.
|
|
||||||
allowed, _ := util.AllowedCPUs()
|
|
||||||
validCPU := 0
|
|
||||||
if len(allowed) > 0 {
|
|
||||||
validCPU = allowed[0]
|
|
||||||
}
|
|
||||||
|
|
||||||
// valid list -> parsed through unchanged
|
|
||||||
assert.Equal(t, []int{validCPU, validCPU}, parseCpuAffinity(newConfig([]any{validCPU, validCPU}), l, 2))
|
|
||||||
|
|
||||||
// a negative entry is out of range on every platform -> disables the override
|
|
||||||
assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, -1}), l, 2))
|
|
||||||
|
|
||||||
// a non-integer entry -> disables the override
|
|
||||||
assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, "not-a-cpu"}), l, 2))
|
|
||||||
|
|
||||||
// a CPU id outside the allowed set -> disables the override. Only assertable
|
|
||||||
// where we can enumerate the allowed set (e.g. linux); 1<<20 is far beyond
|
|
||||||
// any representable CPU id so it can never be in the mask.
|
|
||||||
if len(allowed) > 0 {
|
|
||||||
assert.Nil(t, parseCpuAffinity(newConfig([]any{1 << 20}), l, 1))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+4
-13
@@ -14,8 +14,7 @@ type MessageMetrics struct {
|
|||||||
rxUnknown metrics.Counter
|
rxUnknown metrics.Counter
|
||||||
txUnknown metrics.Counter
|
txUnknown metrics.Counter
|
||||||
|
|
||||||
rxInvalid metrics.Counter
|
rxInvalid metrics.Counter
|
||||||
txExhausted 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) {
|
||||||
@@ -42,13 +41,6 @@ func (m *MessageMetrics) RxInvalid(i int64) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TxExhausted counts outbound packets dropped because the tunnel's message counter is spent.
|
|
||||||
func (m *MessageMetrics) TxExhausted(i int64) {
|
|
||||||
if m != nil && m.txExhausted != nil {
|
|
||||||
m.txExhausted.Inc(i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newMessageMetrics() *MessageMetrics {
|
func newMessageMetrics() *MessageMetrics {
|
||||||
gen := func(t string) [][]metrics.Counter {
|
gen := func(t string) [][]metrics.Counter {
|
||||||
return [][]metrics.Counter{
|
return [][]metrics.Counter{
|
||||||
@@ -69,10 +61,9 @@ func newMessageMetrics() *MessageMetrics {
|
|||||||
rx: gen("rx"),
|
rx: gen("rx"),
|
||||||
tx: gen("tx"),
|
tx: gen("tx"),
|
||||||
|
|
||||||
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),
|
rxInvalid: metrics.GetOrRegisterCounter("messages.rx.invalid", nil),
|
||||||
txExhausted: metrics.GetOrRegisterCounter("messages.tx.exhausted", nil),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -25,9 +25,6 @@ func (s *CipherStateAESGCM) EncryptDanger(out, ad, plaintext []byte, n uint64, n
|
|||||||
if s == nil {
|
if s == nil {
|
||||||
return nil, errors.New("no cipher state available to encrypt")
|
return nil, errors.New("no cipher state available to encrypt")
|
||||||
}
|
}
|
||||||
if n >= RejectAfterMessages {
|
|
||||||
return nil, ErrMessageCounterExhausted
|
|
||||||
}
|
|
||||||
nb[0] = 0
|
nb[0] = 0
|
||||||
nb[1] = 0
|
nb[1] = 0
|
||||||
nb[2] = 0
|
nb[2] = 0
|
||||||
|
|||||||
+65
-4
@@ -4,16 +4,77 @@
|
|||||||
package noiseutil
|
package noiseutil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/boring"
|
"crypto/aes"
|
||||||
|
"crypto/cipher"
|
||||||
|
"encoding/binary"
|
||||||
|
|
||||||
|
// unsafe needed for go:linkname
|
||||||
|
_ "unsafe"
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
"github.com/flynn/noise"
|
||||||
)
|
)
|
||||||
|
|
||||||
var CipherAESGCM noise.CipherFunc = CipherAESGCMFIPS140
|
|
||||||
|
|
||||||
// EncryptLockNeeded indicates if calls to Encrypt need a lock
|
// EncryptLockNeeded indicates if calls to Encrypt need a lock
|
||||||
// This is true for boringcrypto because the Seal function verifies that the
|
// This is true for boringcrypto because the Seal function verifies that the
|
||||||
// nonce is strictly increasing.
|
// nonce is strictly increasing.
|
||||||
const EncryptLockNeeded = true
|
const EncryptLockNeeded = true
|
||||||
|
|
||||||
var boringEnabled = boring.Enabled()
|
// NewGCMTLS is no longer exposed in go1.19+, so we need to link it in
|
||||||
|
// See: https://github.com/golang/go/issues/56326
|
||||||
|
//
|
||||||
|
// NewGCMTLS is the internal method used with boringcrypto that provides a
|
||||||
|
// validated mode of AES-GCM which enforces the nonce is strictly
|
||||||
|
// monotonically increasing. This is the TLS 1.2 specification for nonce
|
||||||
|
// generation (which also matches the method used by the Noise Protocol)
|
||||||
|
//
|
||||||
|
// - https://github.com/golang/go/blob/go1.19/src/crypto/tls/cipher_suites.go#L520-L522
|
||||||
|
// - https://github.com/golang/go/blob/go1.19/src/crypto/internal/boring/aes.go#L235-L237
|
||||||
|
// - https://github.com/golang/go/blob/go1.19/src/crypto/internal/boring/aes.go#L250
|
||||||
|
// - https://github.com/google/boringssl/blob/ae223d6138807a13006342edfeef32e813246b39/include/openssl/aead.h#L379-L381
|
||||||
|
// - https://github.com/google/boringssl/blob/ae223d6138807a13006342edfeef32e813246b39/crypto/fipsmodule/cipher/e_aes.c#L1082-L1093
|
||||||
|
//
|
||||||
|
//go:linkname newGCMTLS crypto/internal/boring.NewGCMTLS
|
||||||
|
func newGCMTLS(c cipher.Block) (cipher.AEAD, error)
|
||||||
|
|
||||||
|
type cipherFn struct {
|
||||||
|
fn func([32]byte) noise.Cipher
|
||||||
|
name string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c cipherFn) Cipher(k [32]byte) noise.Cipher { return c.fn(k) }
|
||||||
|
func (c cipherFn) CipherName() string { return c.name }
|
||||||
|
|
||||||
|
// CipherAESGCM is the AES256-GCM AEAD cipher (using NewGCMTLS when GoBoring is present)
|
||||||
|
var CipherAESGCM noise.CipherFunc = cipherFn{cipherAESGCMBoring, "AESGCM"}
|
||||||
|
|
||||||
|
func cipherAESGCMBoring(k [32]byte) noise.Cipher {
|
||||||
|
c, err := aes.NewCipher(k[:])
|
||||||
|
if err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
gcm, err := newGCMTLS(c)
|
||||||
|
if err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
return aeadCipher{
|
||||||
|
gcm,
|
||||||
|
func(n uint64) []byte {
|
||||||
|
var nonce [12]byte
|
||||||
|
binary.BigEndian.PutUint64(nonce[4:], n)
|
||||||
|
return nonce[:]
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type aeadCipher struct {
|
||||||
|
cipher.AEAD
|
||||||
|
nonce func(uint64) []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c aeadCipher) Encrypt(out []byte, n uint64, ad, plaintext []byte) []byte {
|
||||||
|
return c.Seal(out, c.nonce(n), plaintext, ad)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c aeadCipher) Decrypt(out []byte, n uint64, ad, ciphertext []byte) ([]byte, error) {
|
||||||
|
return c.Open(out, c.nonce(n), ciphertext, ad)
|
||||||
|
}
|
||||||
|
|||||||
@@ -4,6 +4,8 @@
|
|||||||
package noiseutil
|
package noiseutil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"crypto/boring"
|
||||||
|
"encoding/hex"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
@@ -12,3 +14,33 @@ import (
|
|||||||
func TestEncryptLockNeeded(t *testing.T) {
|
func TestEncryptLockNeeded(t *testing.T) {
|
||||||
assert.True(t, EncryptLockNeeded)
|
assert.True(t, EncryptLockNeeded)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Ensure NewGCMTLS validates the nonce is non-repeating
|
||||||
|
func TestNewGCMTLS(t *testing.T) {
|
||||||
|
assert.True(t, boring.Enabled())
|
||||||
|
|
||||||
|
// Test Case 16 from GCM Spec:
|
||||||
|
// - (now dead link): http://csrc.nist.gov/groups/ST/toolkit/BCM/documents/proposedmodes/gcm/gcm-spec.pdf
|
||||||
|
// - as listed in boringssl tests: https://github.com/google/boringssl/blob/fips-20220613/crypto/cipher_extra/test/cipher_tests.txt#L412-L418
|
||||||
|
key, _ := hex.DecodeString("feffe9928665731c6d6a8f9467308308feffe9928665731c6d6a8f9467308308")
|
||||||
|
iv, _ := hex.DecodeString("cafebabefacedbaddecaf888")
|
||||||
|
plaintext, _ := hex.DecodeString("d9313225f88406e5a55909c5aff5269a86a7a9531534f7da2e4c303d8a318a721c3c0c95956809532fcf0e2449a6b525b16aedf5aa0de657ba637b39")
|
||||||
|
aad, _ := hex.DecodeString("feedfacedeadbeeffeedfacedeadbeefabaddad2")
|
||||||
|
expected, _ := hex.DecodeString("522dc1f099567d07f47f37a32a84427d643a8cdcbfe5c0c97598a2bd2555d1aa8cb08e48590dbb3da7b08b1056828838c5f61e6393ba7a0abcc9f662")
|
||||||
|
expectedTag, _ := hex.DecodeString("76fc6ece0f4e1768cddf8853bb2d551b")
|
||||||
|
|
||||||
|
expected = append(expected, expectedTag...)
|
||||||
|
|
||||||
|
var keyArray [32]byte
|
||||||
|
copy(keyArray[:], key)
|
||||||
|
c := CipherAESGCM.Cipher(keyArray)
|
||||||
|
aead := c.(aeadCipher).AEAD
|
||||||
|
|
||||||
|
dst := aead.Seal([]byte{}, iv, plaintext, aad)
|
||||||
|
assert.Equal(t, expected, dst)
|
||||||
|
|
||||||
|
// We expect this to fail since we are re-encrypting with a repeat IV
|
||||||
|
assert.PanicsWithError(t, "boringcrypto: EVP_AEAD_CTX_seal failed", func() {
|
||||||
|
dst = aead.Seal([]byte{}, iv, plaintext, aad)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|||||||
@@ -24,9 +24,6 @@ func (s *CipherStateChaChaPoly) EncryptDanger(out, ad, plaintext []byte, n uint6
|
|||||||
if s == nil {
|
if s == nil {
|
||||||
return nil, errors.New("no cipher state available to encrypt")
|
return nil, errors.New("no cipher state available to encrypt")
|
||||||
}
|
}
|
||||||
if n >= RejectAfterMessages {
|
|
||||||
return nil, ErrMessageCounterExhausted
|
|
||||||
}
|
|
||||||
nb[0] = 0
|
nb[0] = 0
|
||||||
nb[1] = 0
|
nb[1] = 0
|
||||||
nb[2] = 0
|
nb[2] = 0
|
||||||
|
|||||||
@@ -1,22 +1,11 @@
|
|||||||
package noiseutil
|
package noiseutil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"math"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
"github.com/flynn/noise"
|
||||||
)
|
)
|
||||||
|
|
||||||
// RejectHeadroom is the wrap gap for senders racing the counter, sized large enough for any routine count.
|
|
||||||
const RejectHeadroom = uint64(1) << 40
|
|
||||||
|
|
||||||
// RejectAfterMessages is the nonce ceiling: encrypting stops RejectHeadroom short of the wrap.
|
|
||||||
const RejectAfterMessages = math.MaxUint64 - RejectHeadroom
|
|
||||||
|
|
||||||
// ErrMessageCounterExhausted is returned by EncryptDanger once the nonce reaches RejectAfterMessages.
|
|
||||||
var ErrMessageCounterExhausted = errors.New("message counter exhausted")
|
|
||||||
|
|
||||||
// CipherState is the post-handshake AEAD cipher used for the data plane.
|
// CipherState is the post-handshake AEAD cipher used for the data plane.
|
||||||
// Each supported cipher has its own concrete implementation in this package with the nonce endianness hardcoded,
|
// Each supported cipher has its own concrete implementation in this package with the nonce endianness hardcoded,
|
||||||
// so the encrypt/decrypt fast path avoids interface dispatch on the byte order.
|
// so the encrypt/decrypt fast path avoids interface dispatch on the byte order.
|
||||||
@@ -40,11 +29,8 @@ type CipherState interface {
|
|||||||
// NewCipherState wraps the post-handshake noise.CipherState in the per-cipher type that matches cipherFunc.
|
// NewCipherState wraps the post-handshake noise.CipherState in the per-cipher type that matches cipherFunc.
|
||||||
// cipherFunc must be the same cipher used to build the noise CipherSuite that produced s.
|
// cipherFunc must be the same cipher used to build the noise CipherSuite that produced s.
|
||||||
func NewCipherState(s *noise.CipherState, cipherFunc noise.CipherFunc) CipherState {
|
func NewCipherState(s *noise.CipherState, cipherFunc noise.CipherFunc) CipherState {
|
||||||
if cs, ok := s.Cipher().(CipherState); ok {
|
|
||||||
return cs
|
|
||||||
}
|
|
||||||
switch cipherFunc.CipherName() {
|
switch cipherFunc.CipherName() {
|
||||||
case noise.CipherAESGCM.CipherName():
|
case CipherAESGCM.CipherName():
|
||||||
return NewCipherStateAESGCM(s)
|
return NewCipherStateAESGCM(s)
|
||||||
case noise.CipherChaChaPoly.CipherName():
|
case noise.CipherChaChaPoly.CipherName():
|
||||||
return NewCipherStateChaChaPoly(s)
|
return NewCipherStateChaChaPoly(s)
|
||||||
|
|||||||
@@ -1,8 +1,6 @@
|
|||||||
package noiseutil
|
package noiseutil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/fips140"
|
|
||||||
"math"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
"github.com/flynn/noise"
|
||||||
@@ -12,30 +10,24 @@ import (
|
|||||||
|
|
||||||
func TestCipherStateAESGCMRoundtrip(t *testing.T) {
|
func TestCipherStateAESGCMRoundtrip(t *testing.T) {
|
||||||
enc, dec := buildCipherStates(t, CipherAESGCM)
|
enc, dec := buildCipherStates(t, CipherAESGCM)
|
||||||
roundtrip(t, NewCipherState(enc, CipherAESGCM), NewCipherState(dec, CipherAESGCM))
|
roundtrip(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCipherStateChaChaPolyRoundtrip(t *testing.T) {
|
func TestCipherStateChaChaPolyRoundtrip(t *testing.T) {
|
||||||
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
|
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||||
roundtrip(t, NewCipherState(enc, noise.CipherChaChaPoly), NewCipherState(dec, noise.CipherChaChaPoly))
|
roundtrip(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestNewCipherStateDispatch(t *testing.T) {
|
func TestNewCipherStateDispatch(t *testing.T) {
|
||||||
encA, _ := buildCipherStates(t, CipherAESGCM)
|
encA, _ := buildCipherStates(t, CipherAESGCM)
|
||||||
encC, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
encC, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||||
|
|
||||||
if !boringEnabled && !fips140.Enabled() {
|
assert.IsType(t, &CipherStateAESGCM{}, NewCipherState(encA, CipherAESGCM))
|
||||||
assert.IsType(t, &CipherStateAESGCM{}, NewCipherState(encA, CipherAESGCM))
|
|
||||||
} else {
|
|
||||||
// fips140
|
|
||||||
assert.IsType(t, encA.Cipher(), NewCipherState(encA, CipherAESGCM))
|
|
||||||
}
|
|
||||||
|
|
||||||
assert.IsType(t, &CipherStateChaChaPoly{}, NewCipherState(encC, noise.CipherChaChaPoly))
|
assert.IsType(t, &CipherStateChaChaPoly{}, NewCipherState(encC, noise.CipherChaChaPoly))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestNewCipherStateUnsupportedPanics(t *testing.T) {
|
func TestNewCipherStateUnsupportedPanics(t *testing.T) {
|
||||||
enc, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
enc, _ := buildCipherStates(t, CipherAESGCM)
|
||||||
assert.Panics(t, func() {
|
assert.Panics(t, func() {
|
||||||
NewCipherState(enc, fakeCipher{})
|
NewCipherState(enc, fakeCipher{})
|
||||||
})
|
})
|
||||||
@@ -97,24 +89,6 @@ func roundtrip(t *testing.T, enc, dec CipherState) {
|
|||||||
assert.Equal(t, 16, enc.Overhead())
|
assert.Equal(t, 16, enc.Overhead())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestEncryptRejectsExhaustedCounter(t *testing.T) {
|
|
||||||
// Pin the headroom below the uint64 wrap so a typo can't silently move the ceiling.
|
|
||||||
require.Equal(t, uint64(1)<<40, RejectHeadroom)
|
|
||||||
require.Equal(t, math.MaxUint64-RejectHeadroom, RejectAfterMessages)
|
|
||||||
|
|
||||||
encA, _ := buildCipherStates(t, CipherAESGCM)
|
|
||||||
encC, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
|
||||||
nb := make([]byte, 12)
|
|
||||||
|
|
||||||
for _, cs := range []CipherState{NewCipherStateAESGCM(encA), NewCipherStateChaChaPoly(encC)} {
|
|
||||||
_, err := cs.EncryptDanger(nil, nil, []byte("x"), RejectAfterMessages-1, nb)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = cs.EncryptDanger(nil, nil, []byte("x"), RejectAfterMessages, nb)
|
|
||||||
require.ErrorIs(t, err, ErrMessageCounterExhausted)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkCipherStateEncryptAESGCM(b *testing.B) {
|
func BenchmarkCipherStateEncryptAESGCM(b *testing.B) {
|
||||||
enc, _ := buildCipherStatesB(b, CipherAESGCM)
|
enc, _ := buildCipherStatesB(b, CipherAESGCM)
|
||||||
benchEncryptCipherState(b, NewCipherState(enc, CipherAESGCM))
|
benchEncryptCipherState(b, NewCipherState(enc, CipherAESGCM))
|
||||||
@@ -190,48 +164,3 @@ func TestCipherStateNilSafety(t *testing.T) {
|
|||||||
assert.Empty(t, out)
|
assert.Empty(t, out)
|
||||||
assert.Equal(t, 0, cc.Overhead())
|
assert.Equal(t, 0, cc.Overhead())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCipherStateAESGCMInPlaceDecrypt(t *testing.T) {
|
|
||||||
enc, dec := buildCipherStates(t, CipherAESGCM)
|
|
||||||
inPlaceDecrypt(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCipherStateChaChaPolyInPlaceDecrypt(t *testing.T) {
|
|
||||||
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
|
|
||||||
inPlaceDecrypt(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec))
|
|
||||||
}
|
|
||||||
|
|
||||||
func inPlaceDecrypt(t *testing.T, enc, dec CipherState) {
|
|
||||||
t.Helper()
|
|
||||||
const hdrLen = 16
|
|
||||||
plaintext := []byte("in-place decrypt should replace the ciphertext bytes")
|
|
||||||
nb := make([]byte, 12)
|
|
||||||
|
|
||||||
// packet = [16-byte header | ciphertext+tag], like a nebula Message.
|
|
||||||
packet := make([]byte, hdrLen, hdrLen+len(plaintext)+enc.Overhead())
|
|
||||||
for i := range packet {
|
|
||||||
packet[i] = byte(i)
|
|
||||||
}
|
|
||||||
packet, err := enc.EncryptDanger(packet, packet[:hdrLen], plaintext, 1, nb)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Simulate a GRO row: [packet | next segment]. A failed auth on packet
|
|
||||||
// may zero packet's plaintext region but must not touch the header, the
|
|
||||||
// tag, or the neighboring segment.
|
|
||||||
neighbor := []byte("next coalesced segment, must stay intact")
|
|
||||||
row := append(append([]byte(nil), packet...), neighbor...)
|
|
||||||
tampered := row[:len(packet)]
|
|
||||||
tampered[hdrLen] ^= 0x01
|
|
||||||
_, err = dec.DecryptDanger(tampered[hdrLen:hdrLen], tampered[:hdrLen], tampered[hdrLen:], 1, nb)
|
|
||||||
require.Error(t, err)
|
|
||||||
assert.Equal(t, packet[:hdrLen], tampered[:hdrLen], "failed auth must not touch the header")
|
|
||||||
assert.Equal(t, packet[len(packet)-dec.Overhead():], tampered[len(tampered)-dec.Overhead():],
|
|
||||||
"failed auth must not touch the tag")
|
|
||||||
assert.Equal(t, neighbor, row[len(packet):], "failed auth must not touch the next segment")
|
|
||||||
|
|
||||||
out, err := dec.DecryptDanger(packet[hdrLen:hdrLen], packet[:hdrLen], packet[hdrLen:], 1, nb)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, plaintext, out)
|
|
||||||
// The plaintext must be IN the packet buffer, not a fresh allocation.
|
|
||||||
assert.Equal(t, &packet[hdrLen], &out[0], "plaintext must alias the packet buffer")
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,197 +0,0 @@
|
|||||||
package noiseutil
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"crypto/cipher"
|
|
||||||
"crypto/fips140"
|
|
||||||
"encoding/binary"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"reflect"
|
|
||||||
"runtime"
|
|
||||||
"unsafe"
|
|
||||||
|
|
||||||
// unsafe needed for go:linkname
|
|
||||||
_ "crypto/tls"
|
|
||||||
_ "unsafe"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TODO: Use NewGCMWithCounterNonce or NewGCMForQUIC once available:
|
|
||||||
// - https://github.com/golang/go/issues/73110
|
|
||||||
// - https://github.com/golang/go/issues/79219
|
|
||||||
// Using tls.aeadAESGCMTLS13 gives us the TLS 1.3 GCM, which also verifies
|
|
||||||
// that the nonce is strictly increasing. This works for both boringcrypto
|
|
||||||
// and fips140.
|
|
||||||
//
|
|
||||||
//go:linkname aeadAESGCMTLS13 crypto/tls.aeadAESGCMTLS13
|
|
||||||
func aeadAESGCMTLS13(key, noncePrefix []byte) cipher.AEAD
|
|
||||||
|
|
||||||
type cipherFn struct {
|
|
||||||
fn func([32]byte) noise.Cipher
|
|
||||||
name string
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c cipherFn) Cipher(k [32]byte) noise.Cipher { return c.fn(k) }
|
|
||||||
func (c cipherFn) CipherName() string { return c.name }
|
|
||||||
|
|
||||||
// CipherAESGCMFIPS140 is the AES256-GCM AEAD cipher (using tls.aeadAESGCMTLS13, for both boringcrypto and fips140)
|
|
||||||
var CipherAESGCMFIPS140 noise.CipherFunc = cipherFn{cipherAESGCMFIPS140, "AESGCM"}
|
|
||||||
|
|
||||||
// tls.aeadAESGCMTLS13 uses a 4 byte static prefix and an 8 byte XOR mask
|
|
||||||
var emptyNonce = []byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}
|
|
||||||
|
|
||||||
func cipherAESGCMFIPS140(k [32]byte) noise.Cipher {
|
|
||||||
gcm := aeadAESGCMTLS13(k[:], emptyNonce)
|
|
||||||
gcm = extractFIPSAEAD(gcm)
|
|
||||||
return &aeadGCMFIPS140Cipher{
|
|
||||||
AEAD: gcm,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type aeadGCMFIPS140Cipher struct {
|
|
||||||
cipher.AEAD
|
|
||||||
ready bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// Extract the internal FIPS GCM implementation from the tls wrapper. The TLS
|
|
||||||
// wrapper is not thread safe around Open, so instead of locking around it we
|
|
||||||
// can grab the internal implementation that is thread safe. This is the FIPS
|
|
||||||
// module implementation: `crypto/internal/fips140/aes/gcm.GCMWithXORCounterNonce`
|
|
||||||
//
|
|
||||||
// - https://github.com/golang/go/blob/go1.26.4/src/crypto/internal/fips140/aes/gcm/gcm_nonces.go#L212-L287
|
|
||||||
//
|
|
||||||
// The wrapper is struct `crypto/tls.xorNonceAEAD` , with field `aead`:
|
|
||||||
//
|
|
||||||
// - https://github.com/golang/go/blob/go1.26.4/src/crypto/tls/cipher_suites.go#L482-L487
|
|
||||||
//
|
|
||||||
// This can be cleaned up once these FIPS implementations are exposed directly:
|
|
||||||
//
|
|
||||||
// - https://github.com/golang/go/issues/73110
|
|
||||||
func extractFIPSAEAD(xorNonceAEAD cipher.AEAD) cipher.AEAD {
|
|
||||||
r := reflect.ValueOf(xorNonceAEAD)
|
|
||||||
v := r.Elem().FieldByName("aead")
|
|
||||||
if !v.IsValid() {
|
|
||||||
// The internal crypto/tls.xorNonceAEAD struct no longer has an `aead`
|
|
||||||
// field. This can only happen on a Go version this code was not built
|
|
||||||
// against; the package init() self-test guards against ever reaching
|
|
||||||
// this at runtime, so this is a defensive fail-fast.
|
|
||||||
panic(fmt.Sprintf("noiseutil: could not extract FIPS AEAD from %T on %s: no `aead` field (incompatible Go version)", xorNonceAEAD, runtime.Version()))
|
|
||||||
}
|
|
||||||
v2 := reflect.NewAt(v.Type(), unsafe.Pointer(v.UnsafeAddr())).Elem()
|
|
||||||
aead, ok := v2.Interface().(cipher.AEAD)
|
|
||||||
if !ok {
|
|
||||||
panic(fmt.Sprintf("noiseutil: extracted FIPS `aead` field is %s, not a cipher.AEAD, on %s (incompatible Go version)", v2.Type(), runtime.Version()))
|
|
||||||
}
|
|
||||||
return aead
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *aeadGCMFIPS140Cipher) init(nonce []byte) {
|
|
||||||
// GCMWithXORCounterNonce expects that the first call to Seal
|
|
||||||
// is with a counter of `0`, this is how it extracts the nonce mask.
|
|
||||||
// We can clean this up in the future when NewGCMWithCounterNonce or
|
|
||||||
// NewGCMForQUIC are available:
|
|
||||||
if !bytes.Equal(emptyNonce, nonce) {
|
|
||||||
c.AEAD.Seal([]byte{}, emptyNonce, []byte{}, []byte{})
|
|
||||||
}
|
|
||||||
c.ready = true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *aeadGCMFIPS140Cipher) Seal(dst, nonce, plaintext, additionalData []byte) []byte {
|
|
||||||
if !c.ready {
|
|
||||||
c.init(nonce)
|
|
||||||
}
|
|
||||||
return c.AEAD.Seal(dst, nonce, plaintext, additionalData)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *aeadGCMFIPS140Cipher) Encrypt(out []byte, n uint64, ad, plaintext []byte) []byte {
|
|
||||||
return c.Seal(out, aeadGCMFIPS140CipherNonce(n), plaintext, ad)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *aeadGCMFIPS140Cipher) Decrypt(out []byte, n uint64, ad, ciphertext []byte) ([]byte, error) {
|
|
||||||
return c.Open(out, aeadGCMFIPS140CipherNonce(n), ciphertext, ad)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *aeadGCMFIPS140Cipher) EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error) {
|
|
||||||
if c == nil {
|
|
||||||
return nil, errors.New("no cipher state available to encrypt")
|
|
||||||
}
|
|
||||||
if n >= RejectAfterMessages {
|
|
||||||
return nil, ErrMessageCounterExhausted
|
|
||||||
}
|
|
||||||
binary.BigEndian.PutUint64(nb[4:], n)
|
|
||||||
out = c.Seal(out, nb, plaintext, ad)
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *aeadGCMFIPS140Cipher) DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error) {
|
|
||||||
if c == nil {
|
|
||||||
return []byte{}, nil
|
|
||||||
}
|
|
||||||
binary.BigEndian.PutUint64(nb[4:], n)
|
|
||||||
return c.Open(out, nb, ciphertext, ad)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *aeadGCMFIPS140Cipher) Overhead() int {
|
|
||||||
if c == nil {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
return c.AEAD.Overhead()
|
|
||||||
}
|
|
||||||
|
|
||||||
func aeadGCMFIPS140CipherNonce(n uint64) []byte {
|
|
||||||
// GCMWithXORCounterNonce uses a 4 byte static prefix and an 8 byte nonce
|
|
||||||
var nonce [12]byte
|
|
||||||
binary.BigEndian.PutUint64(nonce[4:], n)
|
|
||||||
return nonce[:]
|
|
||||||
}
|
|
||||||
|
|
||||||
func init() {
|
|
||||||
if boringEnabled || fips140.Enabled() {
|
|
||||||
initSelfTestAESGCMFIPS140()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// validates the go:linkname + reflection extraction and the nonce-reuse
|
|
||||||
// protection at startup. cipherAESGCMFIPS140 relies on unexported
|
|
||||||
// crypto/tls and crypto/internal/fips140 internals; if a future Go version changes
|
|
||||||
// those, this fails fast with a clear message instead of panicking per-handshake
|
|
||||||
// (or, worse, silently losing the strictly-increasing nonce check that is the whole
|
|
||||||
// point of using this cipher).
|
|
||||||
func initSelfTestAESGCMFIPS140() {
|
|
||||||
var key [32]byte
|
|
||||||
c := cipherAESGCMFIPS140(key)
|
|
||||||
|
|
||||||
// Verify the extracted AEAD produces a working encrypt/decrypt roundtrip.
|
|
||||||
plaintext := []byte("nebula fips140 self-test")
|
|
||||||
ad := []byte("ad")
|
|
||||||
ct := c.Encrypt(nil, 1, ad, plaintext)
|
|
||||||
pt, err := c.Decrypt(nil, 1, ad, ct)
|
|
||||||
if err != nil {
|
|
||||||
panic(fmt.Sprintf("noiseutil: FIPS AES-GCM self-test roundtrip failed on %s: %v", runtime.Version(), err))
|
|
||||||
}
|
|
||||||
if !bytes.Equal(pt, plaintext) {
|
|
||||||
panic(fmt.Sprintf("noiseutil: FIPS AES-GCM self-test roundtrip returned wrong plaintext on %s", runtime.Version()))
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify the nonce-reuse protection still fires: re-encrypting with the same
|
|
||||||
// counter must panic. This is the defensive check that FIPS-140 requires, so
|
|
||||||
// if the extraction ever silently yields an AEAD without it, refuse to start.
|
|
||||||
if !reusePanics(c) {
|
|
||||||
panic(fmt.Sprintf("noiseutil: FIPS AES-GCM self-test did not reject a reused nonce on %s; nonce-reuse protection is missing (incompatible Go version)", runtime.Version()))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// reusePanics reports whether re-encrypting with an already-used counter panics,
|
|
||||||
// as GCMWithXORCounterNonce is expected to.
|
|
||||||
func reusePanics(c noise.Cipher) (panicked bool) {
|
|
||||||
c.Encrypt(nil, 2, nil, nil)
|
|
||||||
defer func() {
|
|
||||||
if recover() != nil {
|
|
||||||
panicked = true
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
c.Encrypt(nil, 2, nil, nil)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
@@ -1,48 +0,0 @@
|
|||||||
package noiseutil
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/cipher"
|
|
||||||
"crypto/fips140"
|
|
||||||
"encoding/hex"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Ensure NewAESGCM validates the nonce is non-repeating
|
|
||||||
func TestNewAESGCM(t *testing.T) {
|
|
||||||
if !boringEnabled && !fips140.Enabled() {
|
|
||||||
t.Skip("TestNewAESGCM is only for fips140/boringcrypto")
|
|
||||||
}
|
|
||||||
|
|
||||||
key, _ := hex.DecodeString("feffe9928665731c6d6a8f9467308308feffe9928665731c6d6a8f9467308308")
|
|
||||||
iv, _ := hex.DecodeString("00000000facedbaddecaf888")
|
|
||||||
plaintext, _ := hex.DecodeString("d9313225f88406e5a55909c5aff5269a86a7a9531534f7da2e4c303d8a318a721c3c0c95956809532fcf0e2449a6b525b16aedf5aa0de657ba637b39")
|
|
||||||
aad, _ := hex.DecodeString("feedfacedeadbeeffeedfacedeadbeefabaddad2")
|
|
||||||
expected, _ := hex.DecodeString("6a65c2edd45bd63c7e29f40e3d2ed8ba2b99f4c83135383d5676652f255059ceb24863ff10afb1089db701245da87fb88d3acd5f9dd0770cac220c3c04145caf25e190aeb775e7080401c628")
|
|
||||||
|
|
||||||
var keyArray [32]byte
|
|
||||||
copy(keyArray[:], key)
|
|
||||||
c := CipherAESGCM.Cipher(keyArray)
|
|
||||||
aead := c.(cipher.AEAD)
|
|
||||||
|
|
||||||
dst := aead.Seal([]byte{}, iv, plaintext, aad)
|
|
||||||
t.Logf("%x", dst)
|
|
||||||
assert.Equal(t, expected, dst)
|
|
||||||
|
|
||||||
// We expect this to fail since we are re-encrypting with a repeat IV
|
|
||||||
switch {
|
|
||||||
case boringEnabled:
|
|
||||||
assert.PanicsWithError(t, "boringcrypto: EVP_AEAD_CTX_seal failed", func() {
|
|
||||||
dst = aead.Seal([]byte{}, iv, plaintext, aad)
|
|
||||||
})
|
|
||||||
case fips140.Version() == "v1.0.0":
|
|
||||||
assert.PanicsWithValue(t, "crypto/cipher: counter decreased", func() {
|
|
||||||
dst = aead.Seal([]byte{}, iv, plaintext, aad)
|
|
||||||
})
|
|
||||||
default:
|
|
||||||
assert.PanicsWithValue(t, "crypto/cipher: counter decreased or remained the same", func() {
|
|
||||||
dst = aead.Seal([]byte{}, iv, plaintext, aad)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,13 +0,0 @@
|
|||||||
//go:build fips140enforce
|
|
||||||
|
|
||||||
package noiseutil
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/fips140"
|
|
||||||
)
|
|
||||||
|
|
||||||
func init() {
|
|
||||||
if !fips140.Enforced() {
|
|
||||||
panic("Nebula compiled with fips140 expects FIPS140 to be enforced. Do not set GODEBUG=fips140, or if you do it must be set as GODEBUG=fips140=only")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+4
-15
@@ -1,25 +1,14 @@
|
|||||||
//go:build !boringcrypto
|
//go:build !boringcrypto
|
||||||
|
// +build !boringcrypto
|
||||||
|
|
||||||
package noiseutil
|
package noiseutil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/fips140"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
"github.com/flynn/noise"
|
||||||
)
|
)
|
||||||
|
|
||||||
// EncryptLockNeeded indicates if calls to Encrypt need a lock
|
// EncryptLockNeeded indicates if calls to Encrypt need a lock
|
||||||
var EncryptLockNeeded = fips140.Enabled()
|
const EncryptLockNeeded = false
|
||||||
|
|
||||||
var CipherAESGCM noise.CipherFunc = initAESGCM()
|
// CipherAESGCM is the standard noise.CipherAESGCM when boringcrypto is not enabled
|
||||||
|
var CipherAESGCM noise.CipherFunc = noise.CipherAESGCM
|
||||||
func initAESGCM() noise.CipherFunc {
|
|
||||||
if fips140.Enabled() {
|
|
||||||
return CipherAESGCMFIPS140
|
|
||||||
} else {
|
|
||||||
return noise.CipherAESGCM
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
var boringEnabled = false
|
|
||||||
|
|||||||
@@ -0,0 +1,14 @@
|
|||||||
|
//go:build !boringcrypto
|
||||||
|
// +build !boringcrypto
|
||||||
|
|
||||||
|
package noiseutil
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestEncryptLockNeeded(t *testing.T) {
|
||||||
|
assert.False(t, EncryptLockNeeded)
|
||||||
|
}
|
||||||
+122
-95
@@ -8,12 +8,11 @@ import (
|
|||||||
"net/netip"
|
"net/netip"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/gopacket/layers"
|
||||||
"golang.org/x/net/ipv6"
|
"golang.org/x/net/ipv6"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/iputil"
|
|
||||||
"github.com/slackhq/nebula/overlay/batch"
|
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -23,11 +22,7 @@ const (
|
|||||||
|
|
||||||
var ErrOutOfWindow = errors.New("out of window packet")
|
var ErrOutOfWindow = errors.New("out of window packet")
|
||||||
|
|
||||||
// readOutsidePackets processes one received underlay packet.
|
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
// Message payloads are decrypted IN PLACE, so packet must stay untouched
|
|
||||||
// by the caller until the batcher for queue q has been flushed
|
|
||||||
func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxContext) {
|
|
||||||
h := rxc.h
|
|
||||||
err := h.Parse(packet)
|
err := h.Parse(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
||||||
@@ -95,7 +90,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
|
|||||||
if isMessageRelay {
|
if isMessageRelay {
|
||||||
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
|
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
|
||||||
} else {
|
} else {
|
||||||
hostinfo = f.hostMap.QueryIndexCached(h.RemoteIndex, rxc.hostmapCache)
|
hostinfo = f.hostMap.QueryIndex(h.RemoteIndex)
|
||||||
}
|
}
|
||||||
|
|
||||||
// At this point we should have a valid existing tunnel, verify and send
|
// At this point we should have a valid existing tunnel, verify and send
|
||||||
@@ -118,18 +113,17 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
|
|||||||
// All remaining packets are encrypted
|
// All remaining packets are encrypted
|
||||||
if isMessageRelay {
|
if isMessageRelay {
|
||||||
// Relay packets are special, this branch should always early-return
|
// Relay packets are special, this branch should always early-return
|
||||||
err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, rxc.nb)
|
if err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, nb); err != nil {
|
||||||
if err != nil {
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h)
|
hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
f.handleOutsideRelayPacket(hostinfo, via, packet, rxc)
|
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
out, err := hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, packet, rxc.nb)
|
out, err = hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, out, packet, nb)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("Failed to decrypt packet", "error", err, "from", via, "header", h)
|
hostinfo.logger(f.l).Debug("Failed to decrypt packet", "error", err, "from", via, "header", h)
|
||||||
@@ -145,7 +139,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
|
|||||||
case header.Message:
|
case header.Message:
|
||||||
switch h.Subtype {
|
switch h.Subtype {
|
||||||
case header.MessageNone:
|
case header.MessageNone:
|
||||||
f.handleOutsideMessagePacket(hostinfo, h.MessageCounter, out, rxc)
|
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache)
|
||||||
default:
|
default:
|
||||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
||||||
return
|
return
|
||||||
@@ -153,23 +147,15 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
|
|||||||
|
|
||||||
case header.LightHouse:
|
case header.LightHouse:
|
||||||
//TODO: assert via is not relayed
|
//TODO: assert via is not relayed
|
||||||
rxc.lhh.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
||||||
|
|
||||||
case header.Test:
|
case header.Test:
|
||||||
switch h.Subtype {
|
switch h.Subtype {
|
||||||
case header.TestReply:
|
case header.TestReply:
|
||||||
// No-op, useful for the Roaming and connectionManager side-effects above
|
// No-op, useful for the Roaming and connectionManager side-effects above
|
||||||
case header.TestRequest:
|
case header.TestRequest:
|
||||||
const maxCipherOverhead = 16 //todo we use this too often, needs a real importable const
|
//recycle the input packet ciphertext as our output buffer
|
||||||
const maxOverhead = header.Len + header.Len + maxCipherOverhead + maxCipherOverhead
|
f.send(header.Test, header.TestReply, hostinfo.ConnectionState, hostinfo, out, nb, packet)
|
||||||
if maxOverhead+len(out) > len(rxc.scratch) {
|
|
||||||
// A reply that cannot fit in scratch is dropped no matter the log level.
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hostinfo.logger(f.l).Debug("dropping oversized test request", "payloadLen", len(out), "from", via)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
f.send(header.Test, header.TestReply, hostinfo.ConnectionState, hostinfo, out, rxc.nb, rxc.scratch[:0])
|
|
||||||
default:
|
default:
|
||||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
|
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
|
||||||
return
|
return
|
||||||
@@ -187,8 +173,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, packet []byte, rxc *rxContext) {
|
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
h := rxc.h
|
|
||||||
// Successfully validated the thing. Get rid of the Relay header and the AEAD tag
|
// Successfully validated the thing. Get rid of the Relay header and the AEAD tag
|
||||||
signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||||
// Pull the Roaming parts up here, and return in all call paths.
|
// Pull the Roaming parts up here, and return in all call paths.
|
||||||
@@ -201,7 +186,9 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
if !ok {
|
if !ok {
|
||||||
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
|
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
|
||||||
// its internal mapping. This should never happen.
|
// its internal mapping. This should never happen.
|
||||||
hostinfo.logger(f.l).Error("HostInfo missing remote relay index", "relayRemoteIndex", h.RemoteIndex)
|
hostinfo.logger(f.l).Error("HostInfo missing remote relay index",
|
||||||
|
"relayRemoteIndex", h.RemoteIndex,
|
||||||
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -215,7 +202,7 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
relay: relay,
|
relay: relay,
|
||||||
IsRelayed: true,
|
IsRelayed: true,
|
||||||
}
|
}
|
||||||
f.readOutsidePackets(via, signedPayload, rxc)
|
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache)
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Find the target HostInfo relay object
|
// Find the target HostInfo relay object
|
||||||
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
||||||
@@ -234,9 +221,8 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Forward this packet through the relay tunnel, rebuilding it in place.
|
// Forward this packet through the relay tunnel, rebuilding it in place.
|
||||||
// Encode overwrites the old outer header, and the new AEAD tag lands where the old one was
|
// Encode overwrites the old outer header, and the new AEAD tag lands where the old one was
|
||||||
fwdBuf := packet[:0]
|
fwdBuf := packet[:0:len(packet)] // Cap to len(packet) to protect memory from a larger parent buffer
|
||||||
//todo it would potentially be nice to batch these
|
f.SendVia(targetHI, targetRelay, signedPayload, nb, fwdBuf, true)
|
||||||
f.SendVia(targetHI, targetRelay, signedPayload, rxc.nb, fwdBuf, true, rxc.q)
|
|
||||||
case TerminalType:
|
case TerminalType:
|
||||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||||
return
|
return
|
||||||
@@ -313,14 +299,11 @@ var (
|
|||||||
ErrIPv4InvalidHeaderLength = errors.New("invalid ipv4 header length")
|
ErrIPv4InvalidHeaderLength = errors.New("invalid ipv4 header length")
|
||||||
ErrIPv4PacketTooShort = errors.New("ipv4 packet is too short")
|
ErrIPv4PacketTooShort = errors.New("ipv4 packet is too short")
|
||||||
ErrIPv6PacketTooShort = errors.New("ipv6 packet is too short")
|
ErrIPv6PacketTooShort = errors.New("ipv6 packet is too short")
|
||||||
|
ErrIPv6CouldNotFindPayload = errors.New("could not find payload in ipv6 packet")
|
||||||
)
|
)
|
||||||
|
|
||||||
// newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers
|
// newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers
|
||||||
func newPacket(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||||
// fp is reused across packets; reset the parse byproducts so an early-error return cannot
|
|
||||||
// leak the previous packet's offsets.
|
|
||||||
fp.IPHdrLen = 0
|
|
||||||
fp.FragAny = false
|
|
||||||
if len(data) < 1 {
|
if len(data) < 1 {
|
||||||
return ErrPacketTooShort
|
return ErrPacketTooShort
|
||||||
}
|
}
|
||||||
@@ -335,7 +318,7 @@ func newPacket(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
|||||||
return ErrUnknownIPVersion
|
return ErrUnknownIPVersion
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||||
dataLen := len(data)
|
dataLen := len(data)
|
||||||
if dataLen < ipv6.HeaderLen {
|
if dataLen < ipv6.HeaderLen {
|
||||||
return ErrIPv6PacketTooShort
|
return ErrIPv6PacketTooShort
|
||||||
@@ -349,64 +332,104 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
|||||||
fp.RemoteAddr, _ = netip.AddrFromSlice(data[24:40])
|
fp.RemoteAddr, _ = netip.AddrFromSlice(data[24:40])
|
||||||
}
|
}
|
||||||
|
|
||||||
// Walk the extension header chain to the upper layer protocol. iputil.IPv6FindUpperProtocol is the single
|
protoAt := 6 // NextHeader is at 6 bytes into the ipv6 header
|
||||||
// source of truth for which headers are extension headers, so this stays in lockstep with the reject path
|
offset := ipv6.HeaderLen // Start at the end of the ipv6 header
|
||||||
// and cannot drift into misreading an unknown protocol (SCTP, GRE, etc.) as a forged transport.
|
next := 0
|
||||||
proto, offset, isFragment, anyFragment, err := iputil.IPv6FindUpperProtocol(data)
|
for {
|
||||||
if err != nil {
|
if protoAt >= dataLen {
|
||||||
return ErrIPv6PacketTooShort
|
break
|
||||||
}
|
|
||||||
|
|
||||||
fp.Protocol = proto
|
|
||||||
fp.Fragment = isFragment
|
|
||||||
fp.FragAny = anyFragment
|
|
||||||
fp.IPHdrLen = offset
|
|
||||||
if isFragment {
|
|
||||||
// Non-first fragments carry no transport header, so we have no ports to read
|
|
||||||
fp.RemotePort = 0
|
|
||||||
fp.LocalPort = 0
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
switch proto {
|
|
||||||
case iputil.IPProtocolICMPv6:
|
|
||||||
// An ICMPv6 message is at least type, code and checksum, 4 bytes. Only echo carries more than we read.
|
|
||||||
if dataLen < offset+4 {
|
|
||||||
return ErrIPv6PacketTooShort
|
|
||||||
}
|
}
|
||||||
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
|
proto := layers.IPProtocol(data[protoAt])
|
||||||
switch data[offset] { //icmp type
|
|
||||||
case iputil.ICMPv6TypeEchoRequest, iputil.ICMPv6TypeEchoReply:
|
switch proto {
|
||||||
|
case layers.IPProtocolESP, layers.IPProtocolNoNextHeader:
|
||||||
|
fp.Protocol = uint8(proto)
|
||||||
|
fp.RemotePort = 0
|
||||||
|
fp.LocalPort = 0
|
||||||
|
fp.Fragment = false
|
||||||
|
return nil
|
||||||
|
|
||||||
|
case layers.IPProtocolICMPv6:
|
||||||
if dataLen < offset+6 {
|
if dataLen < offset+6 {
|
||||||
return ErrIPv6PacketTooShort
|
return ErrIPv6PacketTooShort
|
||||||
}
|
}
|
||||||
fp.RemotePort = binary.BigEndian.Uint16(data[offset+4 : offset+6]) //identifier
|
fp.Protocol = uint8(proto)
|
||||||
|
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
|
||||||
|
icmptype := data[offset+1]
|
||||||
|
switch icmptype {
|
||||||
|
case layers.ICMPv6TypeEchoRequest, layers.ICMPv6TypeEchoReply:
|
||||||
|
fp.RemotePort = binary.BigEndian.Uint16(data[offset+4 : offset+6]) //identifier
|
||||||
|
default:
|
||||||
|
fp.RemotePort = 0
|
||||||
|
}
|
||||||
|
fp.Fragment = false
|
||||||
|
return nil
|
||||||
|
|
||||||
|
case layers.IPProtocolTCP, layers.IPProtocolUDP:
|
||||||
|
if dataLen < offset+4 {
|
||||||
|
return ErrIPv6PacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
fp.Protocol = uint8(proto)
|
||||||
|
if incoming {
|
||||||
|
fp.RemotePort = binary.BigEndian.Uint16(data[offset : offset+2])
|
||||||
|
fp.LocalPort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
|
||||||
|
} else {
|
||||||
|
fp.LocalPort = binary.BigEndian.Uint16(data[offset : offset+2])
|
||||||
|
fp.RemotePort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
|
||||||
|
}
|
||||||
|
|
||||||
|
fp.Fragment = false
|
||||||
|
return nil
|
||||||
|
|
||||||
|
case layers.IPProtocolIPv6Fragment:
|
||||||
|
// Fragment header is 8 bytes, need at least offset+4 to read the offset field
|
||||||
|
if dataLen < offset+8 {
|
||||||
|
return ErrIPv6PacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this is the first fragment
|
||||||
|
fragmentOffset := binary.BigEndian.Uint16(data[offset+2:offset+4]) &^ uint16(0x7) // Remove the reserved and M flag bits
|
||||||
|
if fragmentOffset != 0 {
|
||||||
|
// Non-first fragment, use what we have now and stop processing
|
||||||
|
fp.Protocol = data[offset]
|
||||||
|
fp.Fragment = true
|
||||||
|
fp.RemotePort = 0
|
||||||
|
fp.LocalPort = 0
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// The next loop should be the transport layer since we are the first fragment
|
||||||
|
next = 8 // Fragment headers are always 8 bytes
|
||||||
|
|
||||||
|
case layers.IPProtocolAH:
|
||||||
|
// Auth headers, used by IPSec, have a different meaning for header length
|
||||||
|
if dataLen <= offset+1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
next = (int(data[offset+1]) + 2) << 2
|
||||||
|
|
||||||
default:
|
default:
|
||||||
fp.RemotePort = 0
|
// Normal ipv6 header length processing
|
||||||
|
if dataLen <= offset+1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
next = (int(data[offset+1]) + 1) << 3
|
||||||
}
|
}
|
||||||
|
|
||||||
case iputil.IPProtocolTCP, iputil.IPProtocolUDP:
|
if next <= 0 {
|
||||||
if dataLen < offset+4 {
|
// Safety check, each ipv6 header has to be at least 8 bytes
|
||||||
return ErrIPv6PacketTooShort
|
next = 8
|
||||||
}
|
|
||||||
if incoming {
|
|
||||||
fp.RemotePort = binary.BigEndian.Uint16(data[offset : offset+2])
|
|
||||||
fp.LocalPort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
|
|
||||||
} else {
|
|
||||||
fp.LocalPort = binary.BigEndian.Uint16(data[offset : offset+2])
|
|
||||||
fp.RemotePort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
|
|
||||||
}
|
}
|
||||||
|
|
||||||
default:
|
protoAt = offset
|
||||||
// don't set ports for protocols Nebula doesn't inspect
|
offset = offset + next
|
||||||
fp.RemotePort = 0
|
|
||||||
fp.LocalPort = 0
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return ErrIPv6CouldNotFindPayload
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||||
// Do we at least have an ipv4 header worth of data?
|
// Do we at least have an ipv4 header worth of data?
|
||||||
if len(data) < ipv4.HeaderLen {
|
if len(data) < ipv4.HeaderLen {
|
||||||
return ErrIPv4PacketTooShort
|
return ErrIPv4PacketTooShort
|
||||||
@@ -423,10 +446,6 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
|||||||
// Check if this is the second or further fragment of a fragmented packet.
|
// Check if this is the second or further fragment of a fragmented packet.
|
||||||
flagsfrags := binary.BigEndian.Uint16(data[6:8])
|
flagsfrags := binary.BigEndian.Uint16(data[6:8])
|
||||||
fp.Fragment = (flagsfrags & 0x1FFF) != 0
|
fp.Fragment = (flagsfrags & 0x1FFF) != 0
|
||||||
// Any fragmentation at all (MF or offset): first fragments have readable ports for the
|
|
||||||
// firewall but must never be coalesced.
|
|
||||||
fp.FragAny = (flagsfrags & 0x3fff) != 0
|
|
||||||
fp.IPHdrLen = ihl
|
|
||||||
|
|
||||||
// Firewall handles protocol checks
|
// Firewall handles protocol checks
|
||||||
fp.Protocol = data[9]
|
fp.Protocol = data[9]
|
||||||
@@ -434,7 +453,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
|||||||
// Accounting for a variable header length, do we have enough data for our src/dst tuples?
|
// Accounting for a variable header length, do we have enough data for our src/dst tuples?
|
||||||
minLen := ihl
|
minLen := ihl
|
||||||
if !fp.Fragment {
|
if !fp.Fragment {
|
||||||
if fp.Protocol == iputil.IPProtocolICMP {
|
if fp.Protocol == firewall.ProtoICMP {
|
||||||
minLen += minFwPacketLen + 2
|
minLen += minFwPacketLen + 2
|
||||||
} else {
|
} else {
|
||||||
minLen += minFwPacketLen
|
minLen += minFwPacketLen
|
||||||
@@ -456,7 +475,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
|||||||
if fp.Fragment {
|
if fp.Fragment {
|
||||||
fp.RemotePort = 0
|
fp.RemotePort = 0
|
||||||
fp.LocalPort = 0
|
fp.LocalPort = 0
|
||||||
} else if fp.Protocol == iputil.IPProtocolICMP { //note that orientation doesn't matter on ICMP
|
} else if fp.Protocol == firewall.ProtoICMP { //note that orientation doesn't matter on ICMP
|
||||||
fp.RemotePort = binary.BigEndian.Uint16(data[ihl+4 : ihl+6]) //identifier
|
fp.RemotePort = binary.BigEndian.Uint16(data[ihl+4 : ihl+6]) //identifier
|
||||||
fp.LocalPort = 0 //code would be uint16(data[ihl+1])
|
fp.LocalPort = 0 //code would be uint16(data[ihl+1])
|
||||||
} else if incoming {
|
} else if incoming {
|
||||||
@@ -470,23 +489,31 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounter uint64, out []byte, rxc *rxContext) {
|
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
err := newPacket(out, true, rxc.fwPacket)
|
err := newPacket(out, true, fwPacket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet", "error", err, "packet", out)
|
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
||||||
|
"error", err,
|
||||||
|
"packet", out,
|
||||||
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
dropReason := f.firewall.Drop(rxc.fwPacket.Packet, true, hostinfo, f.pki.GetCAPool(), rxc.ctCache.Get())
|
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
if dropReason != nil {
|
if dropReason != nil {
|
||||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, rxc.nb, rxc.scratch, rxc.q)
|
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
|
||||||
|
// This gives us a buffer to build the reject packet in
|
||||||
|
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping inbound packet", "fwPacket", rxc.fwPacket, "reason", dropReason)
|
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
||||||
|
"fwPacket", fwPacket,
|
||||||
|
"reason", dropReason,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
err = f.batchers[rxc.q].Commit(out, batch.SortKey{Epoch: hostinfo.ConnectionState.epoch, Counter: messageCounter}, rxc.fwPacket)
|
_, err = f.readers[q].Write(out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to write to tun", "error", err)
|
f.l.Error("Failed to write to tun", "error", err)
|
||||||
}
|
}
|
||||||
|
|||||||
+24
-190
@@ -9,17 +9,15 @@ import (
|
|||||||
|
|
||||||
"github.com/google/gopacket"
|
"github.com/google/gopacket"
|
||||||
"github.com/google/gopacket/layers"
|
"github.com/google/gopacket/layers"
|
||||||
"github.com/slackhq/nebula/iputil"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
"golang.org/x/net/ipv6"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func Test_newPacket(t *testing.T) {
|
func Test_newPacket(t *testing.T) {
|
||||||
p := &firewall.ParsedPacket{}
|
p := &firewall.Packet{}
|
||||||
|
|
||||||
// length fails
|
// length fails
|
||||||
err := newPacket([]byte{}, true, p)
|
err := newPacket([]byte{}, true, p)
|
||||||
@@ -59,7 +57,7 @@ func Test_newPacket(t *testing.T) {
|
|||||||
Src: net.IPv4(10, 0, 0, 1),
|
Src: net.IPv4(10, 0, 0, 1),
|
||||||
Dst: net.IPv4(10, 0, 0, 2),
|
Dst: net.IPv4(10, 0, 0, 2),
|
||||||
Options: []byte{0, 1, 0, 2},
|
Options: []byte{0, 1, 0, 2},
|
||||||
Protocol: iputil.IPProtocolTCP,
|
Protocol: firewall.ProtoTCP,
|
||||||
}
|
}
|
||||||
|
|
||||||
b, _ = h.Marshal()
|
b, _ = h.Marshal()
|
||||||
@@ -67,7 +65,7 @@ func Test_newPacket(t *testing.T) {
|
|||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p)
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("10.0.0.2"), p.LocalAddr)
|
assert.Equal(t, netip.MustParseAddr("10.0.0.2"), p.LocalAddr)
|
||||||
assert.Equal(t, netip.MustParseAddr("10.0.0.1"), p.RemoteAddr)
|
assert.Equal(t, netip.MustParseAddr("10.0.0.1"), p.RemoteAddr)
|
||||||
assert.Equal(t, uint16(3), p.RemotePort)
|
assert.Equal(t, uint16(3), p.RemotePort)
|
||||||
@@ -98,7 +96,7 @@ func Test_newPacket(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func Test_newPacket_v6(t *testing.T) {
|
func Test_newPacket_v6(t *testing.T) {
|
||||||
p := &firewall.ParsedPacket{}
|
p := &firewall.Packet{}
|
||||||
|
|
||||||
// invalid ipv6
|
// invalid ipv6
|
||||||
ip := layers.IPv6{
|
ip := layers.IPv6{
|
||||||
@@ -117,12 +115,12 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
err = newPacket(buffer.Bytes(), true, p)
|
err = newPacket(buffer.Bytes(), true, p)
|
||||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
|
|
||||||
// A v6 packet with a hop-by-hop extension
|
// A v6 packet with a hop-by-hop extension
|
||||||
// ICMPv6 Payload (Echo Request)
|
// ICMPv6 Payload (Echo Request)
|
||||||
icmpLayer := layers.ICMPv6{
|
icmpLayer := layers.ICMPv6{
|
||||||
TypeCode: layers.CreateICMPv6TypeCode(layers.ICMPv6TypeEchoRequest, 0),
|
TypeCode: layers.ICMPv6TypeEchoRequest,
|
||||||
}
|
}
|
||||||
// Hop-by-Hop Extension Header
|
// Hop-by-Hop Extension Header
|
||||||
hopOption := layers.IPv6HopByHopOption{}
|
hopOption := layers.IPv6HopByHopOption{}
|
||||||
@@ -151,12 +149,12 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
// A full IPv6 header and 1 byte in the first extension, but missing
|
// A full IPv6 header and 1 byte in the first extension, but missing
|
||||||
// the length byte.
|
// the length byte.
|
||||||
err = newPacket(buffer.Bytes()[:41], true, p)
|
err = newPacket(buffer.Bytes()[:41], true, p)
|
||||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
|
|
||||||
// A full IPv6 header plus 1 full extension, but only 1 byte of the
|
// A full IPv6 header plus 1 full extension, but only 1 byte of the
|
||||||
// next layer, missing length byte
|
// next layer, missing length byte
|
||||||
err = newPacket(buffer.Bytes()[:49], true, p)
|
err = newPacket(buffer.Bytes()[:49], true, p)
|
||||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
err = nil
|
err = nil
|
||||||
|
|
||||||
// A good ICMP packet
|
// A good ICMP packet
|
||||||
@@ -169,7 +167,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
icmp := layers.ICMPv6{
|
icmp := layers.ICMPv6{
|
||||||
TypeCode: layers.CreateICMPv6TypeCode(layers.ICMPv6TypeEchoRequest, 0),
|
TypeCode: layers.ICMPv6TypeEchoRequest,
|
||||||
Checksum: 0x1234,
|
Checksum: 0x1234,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -191,18 +189,6 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
assert.Equal(t, uint16(0), p.LocalPort)
|
assert.Equal(t, uint16(0), p.LocalPort)
|
||||||
assert.False(t, p.Fragment)
|
assert.False(t, p.Fragment)
|
||||||
|
|
||||||
// A minimal 4 byte non-echo ICMPv6 message (type, code, checksum), no identifier to read
|
|
||||||
icmpMin := make([]byte, ipv6.HeaderLen+4)
|
|
||||||
copy(icmpMin, buffer.Bytes()[:ipv6.HeaderLen])
|
|
||||||
icmpMin[6] = byte(layers.IPProtocolICMPv6)
|
|
||||||
icmpMin[ipv6.HeaderLen] = 1 // type 1, destination unreachable, not echo
|
|
||||||
err = newPacket(icmpMin, true, p)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, uint8(layers.IPProtocolICMPv6), p.Protocol)
|
|
||||||
assert.Equal(t, uint16(0), p.RemotePort)
|
|
||||||
assert.Equal(t, uint16(0), p.LocalPort)
|
|
||||||
assert.False(t, p.Fragment)
|
|
||||||
|
|
||||||
// A good ESP packet
|
// A good ESP packet
|
||||||
b := buffer.Bytes()
|
b := buffer.Bytes()
|
||||||
b[6] = byte(layers.IPProtocolESP)
|
b[6] = byte(layers.IPProtocolESP)
|
||||||
@@ -227,20 +213,16 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
assert.Equal(t, uint16(0), p.LocalPort)
|
assert.Equal(t, uint16(0), p.LocalPort)
|
||||||
assert.False(t, p.Fragment)
|
assert.False(t, p.Fragment)
|
||||||
|
|
||||||
// An unknown protocol packet, we don't dissect it so we fail closed on its true protocol with no ports
|
// An unknown protocol packet
|
||||||
b = buffer.Bytes()
|
b = buffer.Bytes()
|
||||||
b[6] = 255 // 255 is a reserved protocol number
|
b[6] = 255 // 255 is a reserved protocol number
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p)
|
||||||
require.NoError(t, err)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
assert.Equal(t, uint8(255), p.Protocol)
|
|
||||||
assert.Equal(t, uint16(0), p.RemotePort)
|
|
||||||
assert.Equal(t, uint16(0), p.LocalPort)
|
|
||||||
assert.False(t, p.Fragment)
|
|
||||||
|
|
||||||
// A good UDP packet
|
// A good UDP packet
|
||||||
ip = layers.IPv6{
|
ip = layers.IPv6{
|
||||||
Version: 6,
|
Version: 6,
|
||||||
NextHeader: iputil.IPProtocolUDP,
|
NextHeader: firewall.ProtoUDP,
|
||||||
HopLimit: 128,
|
HopLimit: 128,
|
||||||
SrcIP: net.IPv6linklocalallrouters,
|
SrcIP: net.IPv6linklocalallrouters,
|
||||||
DstIP: net.IPv6linklocalallnodes,
|
DstIP: net.IPv6linklocalallnodes,
|
||||||
@@ -263,7 +245,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
// incoming
|
// incoming
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
||||||
assert.Equal(t, uint16(36123), p.RemotePort)
|
assert.Equal(t, uint16(36123), p.RemotePort)
|
||||||
@@ -273,7 +255,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
// outgoing
|
// outgoing
|
||||||
err = newPacket(b, false, p)
|
err = newPacket(b, false, p)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
|
||||||
assert.Equal(t, uint16(36123), p.LocalPort)
|
assert.Equal(t, uint16(36123), p.LocalPort)
|
||||||
@@ -290,7 +272,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
// incoming
|
// incoming
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
||||||
assert.Equal(t, uint16(36123), p.RemotePort)
|
assert.Equal(t, uint16(36123), p.RemotePort)
|
||||||
@@ -300,7 +282,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
// outgoing
|
// outgoing
|
||||||
err = newPacket(b, false, p)
|
err = newPacket(b, false, p)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
|
||||||
assert.Equal(t, uint16(36123), p.LocalPort)
|
assert.Equal(t, uint16(36123), p.LocalPort)
|
||||||
@@ -345,25 +327,25 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
|
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
||||||
assert.Equal(t, uint16(36123), p.RemotePort)
|
assert.Equal(t, uint16(36123), p.RemotePort)
|
||||||
assert.Equal(t, uint16(22), p.LocalPort)
|
assert.Equal(t, uint16(22), p.LocalPort)
|
||||||
assert.False(t, p.Fragment)
|
assert.False(t, p.Fragment)
|
||||||
|
|
||||||
// Ensure buffer bounds checking during processing, a truncated AH header can't reach the payload
|
// Ensure buffer bounds checking during processing
|
||||||
err = newPacket(b[:41], true, p)
|
err = newPacket(b[:41], true, p)
|
||||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||||
|
|
||||||
// Invalid AH header
|
// Invalid AH header
|
||||||
b = buffer.Bytes()
|
b = buffer.Bytes()
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p)
|
||||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
}
|
}
|
||||||
|
|
||||||
func Test_newPacket_ipv6Fragment(t *testing.T) {
|
func Test_newPacket_ipv6Fragment(t *testing.T) {
|
||||||
p := &firewall.ParsedPacket{}
|
p := &firewall.Packet{}
|
||||||
|
|
||||||
ip := &layers.IPv6{
|
ip := &layers.IPv6{
|
||||||
Version: 6,
|
Version: 6,
|
||||||
@@ -543,7 +525,7 @@ func BenchmarkParseV6(b *testing.B) {
|
|||||||
secondFrag = append(secondFrag, fragHeader...)
|
secondFrag = append(secondFrag, fragHeader...)
|
||||||
secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
|
secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
|
||||||
|
|
||||||
fp := &firewall.ParsedPacket{}
|
fp := &firewall.Packet{}
|
||||||
|
|
||||||
b.Run("Normal", func(b *testing.B) {
|
b.Run("Normal", func(b *testing.B) {
|
||||||
for i := 0; i < b.N; i++ {
|
for i := 0; i < b.N; i++ {
|
||||||
@@ -667,7 +649,7 @@ func serializeAH(ah *layers.IPSecAH) []byte {
|
|||||||
// host OS parses the real header, a firewall port/proto bypass. The fix makes parseV6 land
|
// host OS parses the real header, a firewall port/proto bypass. The fix makes parseV6 land
|
||||||
// on the same offset the host does.
|
// on the same offset the host does.
|
||||||
func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
||||||
p := &firewall.ParsedPacket{}
|
p := &firewall.Packet{}
|
||||||
|
|
||||||
const (
|
const (
|
||||||
hdrLen = 40 // IPv6 header
|
hdrLen = 40 // IPv6 header
|
||||||
@@ -679,7 +661,7 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
|||||||
pkt := make([]byte, realTCPAt+4)
|
pkt := make([]byte, realTCPAt+4)
|
||||||
pkt[0] = 0x60 // version 6
|
pkt[0] = 0x60 // version 6
|
||||||
pkt[6] = byte(layers.IPProtocolIPv6Destination) // NextHeader -> Destination Options
|
pkt[6] = byte(layers.IPProtocolIPv6Destination) // NextHeader -> Destination Options
|
||||||
pkt[40] = byte(iputil.IPProtocolTCP) // Dest-Options NextHeader -> TCP
|
pkt[40] = byte(firewall.ProtoTCP) // Dest-Options NextHeader -> TCP
|
||||||
pkt[41] = 255 // HdrExtLen = 255
|
pkt[41] = 255 // HdrExtLen = 255
|
||||||
|
|
||||||
// Forged transport header at the pre-fix (wrong) offset: dst port 443.
|
// Forged transport header at the pre-fix (wrong) offset: dst port 443.
|
||||||
@@ -688,156 +670,8 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
|||||||
binary.BigEndian.PutUint16(pkt[realTCPAt+2:realTCPAt+4], 22)
|
binary.BigEndian.PutUint16(pkt[realTCPAt+2:realTCPAt+4], 22)
|
||||||
|
|
||||||
require.NoError(t, newPacket(pkt, true, p))
|
require.NoError(t, newPacket(pkt, true, p))
|
||||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||||
// LocalPort is the destination port for incoming traffic. It must be the real port (22)
|
// LocalPort is the destination port for incoming traffic. It must be the real port (22)
|
||||||
// the host delivers to, not the forged 443 at the overflowed offset.
|
// the host delivers to, not the forged 443 at the overflowed offset.
|
||||||
assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset")
|
assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Test_newPacket_v6ExtHeaderPastBuffer is a regression test for an extension header whose declared length
|
|
||||||
// advances the walk past the end of the packet. The upper layer protocol's header isn't actually present,
|
|
||||||
// so parseV6 must drop the packet rather than classify it as the terminal protocol with no ports.
|
|
||||||
func Test_newPacket_v6ExtHeaderPastBuffer(t *testing.T) {
|
|
||||||
p := &firewall.ParsedPacket{}
|
|
||||||
|
|
||||||
pkt := make([]byte, 48)
|
|
||||||
pkt[0] = 0x60
|
|
||||||
pkt[6] = byte(layers.IPProtocolIPv6Destination) // Destination Options
|
|
||||||
pkt[7] = 64 // hop limit
|
|
||||||
pkt[40] = byte(layers.IPProtocolSCTP) // Dest Options next header = SCTP
|
|
||||||
pkt[41] = 255 // declared length (255+1)*8 = 2048, past the 48 byte buffer
|
|
||||||
|
|
||||||
require.ErrorIs(t, newPacket(pkt, true, p), ErrIPv6PacketTooShort)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Test_newPacket_v6ExtHeaderConfusion is a regression test for parseV6 walking any unrecognized
|
|
||||||
// Next Header as if it were an ipv6 extension header. A real upper layer protocol Nebula doesn't
|
|
||||||
// dissect (SCTP here) is not walkable, so applying the (len+1)*8 formula marched into the SCTP
|
|
||||||
// payload and landed on a byte that looked like UDP, forging a protocol/port pair the firewall
|
|
||||||
// would trust while the host delivered the real SCTP datagram. The fix fails closed: the packet
|
|
||||||
// is classified as its true protocol with no ports, so it only matches an `any` rule.
|
|
||||||
func Test_newPacket_v6ExtHeaderConfusion(t *testing.T) {
|
|
||||||
p := &firewall.ParsedPacket{}
|
|
||||||
|
|
||||||
pkt := make([]byte, 52)
|
|
||||||
pkt[0] = 0x60 // version 6
|
|
||||||
pkt[6] = byte(layers.IPProtocolSCTP) // NextHeader = SCTP, a real protocol, not an extension header
|
|
||||||
pkt[7] = 64 // hop limit
|
|
||||||
|
|
||||||
// Real SCTP header at offset 40. Pre-fix parseV6 walked SCTP as an extension header: byte 41 (0x00, the
|
|
||||||
// low byte of the src port below) was read as the header length, giving next=(0+1)*8=8, which landed the
|
|
||||||
// walk on byte 40 (0x11), misread as NextHeader=UDP, then bytes 48-51 as ports.
|
|
||||||
binary.BigEndian.PutUint16(pkt[40:42], 0x1100) // SCTP src port; byte 40=0x11, byte 41=0x00
|
|
||||||
binary.BigEndian.PutUint16(pkt[42:44], 445) // SCTP dst port, never read by parseV6
|
|
||||||
binary.BigEndian.PutUint16(pkt[48:50], 53) // SCTP checksum bytes, pre-fix forged RemotePort
|
|
||||||
binary.BigEndian.PutUint16(pkt[50:52], 53) // pre-fix forged LocalPort
|
|
||||||
|
|
||||||
require.NoError(t, newPacket(pkt, true, p))
|
|
||||||
assert.Equal(t, uint8(layers.IPProtocolSCTP), p.Protocol, "must classify as the true protocol, not the forged UDP")
|
|
||||||
assert.Equal(t, uint16(0), p.RemotePort)
|
|
||||||
assert.Equal(t, uint16(0), p.LocalPort)
|
|
||||||
assert.False(t, p.Fragment)
|
|
||||||
|
|
||||||
// Same confusion, but the unknown protocol sits after a real extension header. The HopByHop is walked
|
|
||||||
// correctly, then SCTP must still fail closed instead of being walked into its own payload. Protocol is
|
|
||||||
// the only assertion that discriminates the fix here, a regression that walked SCTP would misclassify it.
|
|
||||||
chained := make([]byte, 60)
|
|
||||||
chained[0] = 0x60 // version 6
|
|
||||||
chained[6] = byte(layers.IPProtocolIPv6HopByHop) // NextHeader = HopByHop extension
|
|
||||||
chained[7] = 64 // hop limit
|
|
||||||
chained[40] = byte(layers.IPProtocolSCTP) // HopByHop NextHeader = SCTP
|
|
||||||
chained[41] = 0 // HopByHop length 0 -> 8 bytes, SCTP begins at offset 48
|
|
||||||
binary.BigEndian.PutUint16(chained[48:50], 0x1100) // SCTP src port, pre-fix forged NextHeader/length bait
|
|
||||||
binary.BigEndian.PutUint16(chained[50:52], 445) // SCTP dst port, never read by parseV6
|
|
||||||
|
|
||||||
require.NoError(t, newPacket(chained, true, p))
|
|
||||||
assert.Equal(t, uint8(layers.IPProtocolSCTP), p.Protocol, "must fail closed on the unknown protocol after the extension header")
|
|
||||||
assert.Equal(t, uint16(0), p.RemotePort)
|
|
||||||
assert.Equal(t, uint16(0), p.LocalPort)
|
|
||||||
assert.False(t, p.Fragment)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Test_newPacket_parsedFields pins the ParsedPacket byproducts the RX
|
|
||||||
// batcher consumes: IPHdrLen (the true L4 offset) and FragAny (any fragment
|
|
||||||
// shape at all — unlike Packet.Fragment, which is port-oriented and true
|
|
||||||
// only for non-first fragments).
|
|
||||||
func Test_newPacket_parsedFields(t *testing.T) {
|
|
||||||
p := &firewall.ParsedPacket{}
|
|
||||||
|
|
||||||
// Plain IPv4 TCP, IHL 20: L4 offset 20, no fragment shape.
|
|
||||||
v4 := make([]byte, 28)
|
|
||||||
v4[0] = 0x45
|
|
||||||
v4[9] = iputil.IPProtocolTCP
|
|
||||||
binary.BigEndian.PutUint16(v4[6:8], 0x4000) // DF only
|
|
||||||
require.NoError(t, newPacket(v4, true, p))
|
|
||||||
assert.Equal(t, 20, p.IPHdrLen)
|
|
||||||
assert.False(t, p.FragAny)
|
|
||||||
assert.False(t, p.Fragment)
|
|
||||||
|
|
||||||
// IPv4 first fragment (MF set, offset 0): the firewall can read ports
|
|
||||||
// (Fragment false) but the coalescer must not touch it (FragAny true).
|
|
||||||
ff := make([]byte, 28)
|
|
||||||
ff[0] = 0x45
|
|
||||||
ff[9] = iputil.IPProtocolUDP
|
|
||||||
binary.BigEndian.PutUint16(ff[6:8], 0x2000) // MF, offset 0
|
|
||||||
require.NoError(t, newPacket(ff, true, p))
|
|
||||||
assert.False(t, p.Fragment)
|
|
||||||
assert.True(t, p.FragAny)
|
|
||||||
assert.Equal(t, 20, p.IPHdrLen)
|
|
||||||
|
|
||||||
// IPv4 non-first fragment (nonzero offset): both flags set.
|
|
||||||
nf := make([]byte, 28)
|
|
||||||
nf[0] = 0x45
|
|
||||||
nf[9] = iputil.IPProtocolUDP
|
|
||||||
binary.BigEndian.PutUint16(nf[6:8], 0x00b9)
|
|
||||||
require.NoError(t, newPacket(nf, true, p))
|
|
||||||
assert.True(t, p.Fragment)
|
|
||||||
assert.True(t, p.FragAny)
|
|
||||||
|
|
||||||
// IPv4 with options (IHL 24): IPHdrLen tracks the real L4 offset.
|
|
||||||
opts := make([]byte, 32)
|
|
||||||
opts[0] = 0x46
|
|
||||||
opts[9] = iputil.IPProtocolTCP
|
|
||||||
binary.BigEndian.PutUint16(opts[6:8], 0x4000)
|
|
||||||
require.NoError(t, newPacket(opts, true, p))
|
|
||||||
assert.Equal(t, 24, p.IPHdrLen)
|
|
||||||
assert.False(t, p.FragAny)
|
|
||||||
|
|
||||||
// Plain IPv6 TCP: L4 at 40.
|
|
||||||
v6 := make([]byte, 60)
|
|
||||||
v6[0] = 0x60
|
|
||||||
v6[6] = iputil.IPProtocolTCP
|
|
||||||
require.NoError(t, newPacket(v6, true, p))
|
|
||||||
assert.Equal(t, 40, p.IPHdrLen)
|
|
||||||
assert.False(t, p.FragAny)
|
|
||||||
|
|
||||||
// IPv6 hop-by-hop then TCP: IPHdrLen lands past the extension header.
|
|
||||||
hbh := make([]byte, 60)
|
|
||||||
hbh[0] = 0x60
|
|
||||||
hbh[6] = 0 // hop-by-hop
|
|
||||||
hbh[40] = iputil.IPProtocolTCP
|
|
||||||
hbh[41] = 0 // HdrExtLen 0 -> 8-byte header
|
|
||||||
require.NoError(t, newPacket(hbh, true, p))
|
|
||||||
assert.Equal(t, 48, p.IPHdrLen)
|
|
||||||
assert.False(t, p.FragAny)
|
|
||||||
|
|
||||||
// IPv6 first fragment: terminal proto resolved, FragAny set, Fragment not.
|
|
||||||
f6 := make([]byte, 60)
|
|
||||||
f6[0] = 0x60
|
|
||||||
f6[6] = 44 // fragment extension header
|
|
||||||
f6[40] = iputil.IPProtocolUDP
|
|
||||||
require.NoError(t, newPacket(f6, true, p))
|
|
||||||
assert.True(t, p.FragAny)
|
|
||||||
assert.False(t, p.Fragment)
|
|
||||||
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
|
|
||||||
|
|
||||||
// IPv6 non-first fragment: both set, walk stops at the fragment header.
|
|
||||||
f6n := make([]byte, 60)
|
|
||||||
f6n[0] = 0x60
|
|
||||||
f6n[6] = 44
|
|
||||||
f6n[40] = iputil.IPProtocolUDP
|
|
||||||
binary.BigEndian.PutUint16(f6n[42:44], 0x0008)
|
|
||||||
require.NoError(t, newPacket(f6n, true, p))
|
|
||||||
assert.True(t, p.Fragment)
|
|
||||||
assert.True(t, p.FragAny)
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,187 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"math/rand"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
// The checksum-seeding helpers feed the virtio NEEDS_CSUM contract: the L4
|
|
||||||
// checksum field is pre-loaded with the folded (not inverted) pseudo-header
|
|
||||||
// sum, and the kernel later adds the L4 byte sum and inverts. A wrong seed
|
|
||||||
// produces packets every receiver silently drops, with nothing failing on
|
|
||||||
// our side — so these tests check the helpers against an independent
|
|
||||||
// RFC 1071 reference built from explicit pseudo-header bytes, never against
|
|
||||||
// the production checksum code.
|
|
||||||
|
|
||||||
// refSum accumulates big-endian 16-bit words of b (odd tail zero-padded)
|
|
||||||
// into a wide one's-complement accumulator.
|
|
||||||
func refSum(b []byte) uint64 {
|
|
||||||
var s uint64
|
|
||||||
for i := 0; i+1 < len(b); i += 2 {
|
|
||||||
s += uint64(b[i])<<8 | uint64(b[i+1])
|
|
||||||
}
|
|
||||||
if len(b)%2 == 1 {
|
|
||||||
s += uint64(b[len(b)-1]) << 8
|
|
||||||
}
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
|
|
||||||
// refFold folds a wide one's-complement accumulator to 16 bits.
|
|
||||||
func refFold(s uint64) uint16 {
|
|
||||||
for s>>16 != 0 {
|
|
||||||
s = s&0xffff + s>>16
|
|
||||||
}
|
|
||||||
return uint16(s)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFoldOnceNoInvertEdgeCases(t *testing.T) {
|
|
||||||
cases := []uint32{
|
|
||||||
0, 1, 0xffff,
|
|
||||||
0x10000, // single carry
|
|
||||||
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
|
|
||||||
0xffff0000, // high half only
|
|
||||||
0xfffeffff, // fold yields 0x1fffd: needs a second fold
|
|
||||||
0xffffffff, // worst case
|
|
||||||
0x00010001, // simple two-word
|
|
||||||
}
|
|
||||||
for _, c := range cases {
|
|
||||||
want := refFold(uint64(c))
|
|
||||||
if got := foldOnceNoInvert(c); got != want {
|
|
||||||
t.Errorf("foldOnceNoInvert(%#x) = %#x, want %#x", c, got, want)
|
|
||||||
}
|
|
||||||
// Folding a folded value must be a no-op.
|
|
||||||
if got := foldOnceNoInvert(uint32(foldOnceNoInvert(c))); got != foldOnceNoInvert(c) {
|
|
||||||
t.Errorf("foldOnceNoInvert not idempotent at %#x", c)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPseudoSumIPv4MatchesReference(t *testing.T) {
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
src, dst [4]byte
|
|
||||||
proto byte
|
|
||||||
l4Len int
|
|
||||||
}{
|
|
||||||
{"simple", [4]byte{10, 0, 0, 1}, [4]byte{10, 0, 0, 2}, 6, 20},
|
|
||||||
{"zero-len", [4]byte{192, 168, 1, 1}, [4]byte{192, 168, 1, 2}, 17, 0},
|
|
||||||
{"max-len", [4]byte{1, 2, 3, 4}, [4]byte{5, 6, 7, 8}, 6, 65535},
|
|
||||||
{"carry-heavy", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 254}, 17, 65535},
|
|
||||||
{"broadcastish", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 255}, 255, 65535},
|
|
||||||
}
|
|
||||||
for _, c := range cases {
|
|
||||||
t.Run(c.name, func(t *testing.T) {
|
|
||||||
// RFC 793 pseudo-header: src(4) dst(4) zero(1) proto(1) len(2).
|
|
||||||
ph := make([]byte, 12)
|
|
||||||
copy(ph[0:4], c.src[:])
|
|
||||||
copy(ph[4:8], c.dst[:])
|
|
||||||
ph[9] = c.proto
|
|
||||||
binary.BigEndian.PutUint16(ph[10:12], uint16(c.l4Len))
|
|
||||||
want := refFold(refSum(ph))
|
|
||||||
|
|
||||||
got := foldOnceNoInvert(pseudoSumIPv4(c.src[:], c.dst[:], c.proto, c.l4Len))
|
|
||||||
if got != want {
|
|
||||||
t.Errorf("fold(pseudoSumIPv4) = %#x, want %#x", got, want)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPseudoSumIPv6MatchesReference(t *testing.T) {
|
|
||||||
ones := func(b byte) (a [16]byte) {
|
|
||||||
for i := range a {
|
|
||||||
a[i] = b
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
src, dst [16]byte
|
|
||||||
proto byte
|
|
||||||
l4Len int
|
|
||||||
}{
|
|
||||||
{"simple", [16]byte{0xfe, 0x80, 15: 1}, [16]byte{0xfe, 0x80, 15: 2}, 6, 20},
|
|
||||||
{"zero-len", [16]byte{0x20, 0x01, 15: 9}, [16]byte{0x20, 0x01, 15: 8}, 17, 0},
|
|
||||||
{"max-u16-len", ones(0xff), ones(0xfe), 6, 65535},
|
|
||||||
{"len-past-u16", ones(0xff), ones(0xff), 17, 0x12345}, // exercises the 32-bit split
|
|
||||||
}
|
|
||||||
for _, c := range cases {
|
|
||||||
t.Run(c.name, func(t *testing.T) {
|
|
||||||
// RFC 8200 pseudo-header: src(16) dst(16) len(4) zero(3) next(1).
|
|
||||||
ph := make([]byte, 40)
|
|
||||||
copy(ph[0:16], c.src[:])
|
|
||||||
copy(ph[16:32], c.dst[:])
|
|
||||||
binary.BigEndian.PutUint32(ph[32:36], uint32(c.l4Len))
|
|
||||||
ph[39] = c.proto
|
|
||||||
want := refFold(refSum(ph))
|
|
||||||
|
|
||||||
got := foldOnceNoInvert(pseudoSumIPv6(c.src[:], c.dst[:], c.proto, c.l4Len))
|
|
||||||
if got != want {
|
|
||||||
t.Errorf("fold(pseudoSumIPv6) = %#x, want %#x", got, want)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestIPv4HdrChecksumMatchesReference(t *testing.T) {
|
|
||||||
rng := rand.New(rand.NewSource(0x1791))
|
|
||||||
for _, hdrLen := range []int{20, 24, 40, 60} {
|
|
||||||
for trial := 0; trial < 200; trial++ {
|
|
||||||
hdr := make([]byte, hdrLen)
|
|
||||||
rng.Read(hdr)
|
|
||||||
hdr[0] = 0x40 | byte(hdrLen/4)
|
|
||||||
hdr[10], hdr[11] = 0, 0 // checksum field zeroed, as the contract requires
|
|
||||||
|
|
||||||
want := ^refFold(refSum(hdr))
|
|
||||||
got := ipv4HdrChecksum(hdr)
|
|
||||||
if got != want {
|
|
||||||
t.Fatalf("ipv4HdrChecksum(len=%d trial=%d) = %#x, want %#x", hdrLen, trial, got, want)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Receiver-side property: with the checksum stored, the full
|
|
||||||
// header must sum to all-ones.
|
|
||||||
binary.BigEndian.PutUint16(hdr[10:12], got)
|
|
||||||
if v := refFold(refSum(hdr)); v != 0xffff {
|
|
||||||
t.Fatalf("stored checksum does not validate: full-header fold = %#x", v)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestChecksumSeedReceiverAcceptance is the end-to-end property the helpers
|
|
||||||
// exist for: seed the TCP checksum field with fold(pseudoSum), do what the
|
|
||||||
// kernel's NEEDS_CSUM completion does (one's-complement sum over the L4
|
|
||||||
// bytes including the seed, then invert, then store), and verify the result
|
|
||||||
// the way a receiver does (pseudo-header + L4 must sum to all-ones).
|
|
||||||
func TestChecksumSeedReceiverAcceptance(t *testing.T) {
|
|
||||||
rng := rand.New(rand.NewSource(0x1826))
|
|
||||||
for trial := 0; trial < 200; trial++ {
|
|
||||||
src := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
|
|
||||||
dst := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
|
|
||||||
payLen := rng.Intn(1500)
|
|
||||||
l4 := make([]byte, 20+payLen)
|
|
||||||
rng.Read(l4)
|
|
||||||
|
|
||||||
// Seed exactly as flushSlot does.
|
|
||||||
seed := foldOnceNoInvert(pseudoSumIPv4(src[:], dst[:], 6, len(l4)))
|
|
||||||
binary.BigEndian.PutUint16(l4[16:18], seed)
|
|
||||||
|
|
||||||
// Kernel NEEDS_CSUM completion: sum the L4 region (seed included,
|
|
||||||
// which is equivalent to summing with the field zeroed and folding
|
|
||||||
// the seed in), invert, store.
|
|
||||||
final := ^refFold(refSum(l4[:16]) + uint64(seed) + refSum(l4[18:]))
|
|
||||||
binary.BigEndian.PutUint16(l4[16:18], final)
|
|
||||||
|
|
||||||
// Receiver validation.
|
|
||||||
ph := make([]byte, 12)
|
|
||||||
copy(ph[0:4], src[:])
|
|
||||||
copy(ph[4:8], dst[:])
|
|
||||||
ph[9] = 6
|
|
||||||
binary.BigEndian.PutUint16(ph[10:12], uint16(len(l4)))
|
|
||||||
if v := refFold(refSum(ph) + refSum(l4)); v != 0xffff {
|
|
||||||
t.Fatalf("trial %d: receiver rejects packet: fold = %#x (seed=%#x final=%#x payLen=%d)",
|
|
||||||
trial, v, seed, final, payLen)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,169 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"encoding/binary"
|
|
||||||
)
|
|
||||||
|
|
||||||
// SortKey identifies a packet's position in its sender's transmission order.
|
|
||||||
type SortKey struct {
|
|
||||||
// Epoch is a receiver-local ordinal for the tunnel (ConnectionState) that decrypted the packet:
|
|
||||||
// a re-handshake replaces the tunnel outright and the replacement's epoch is higher,
|
|
||||||
// so the old tunnel's packets sort first during the cutover overlap.
|
|
||||||
Epoch uint64
|
|
||||||
// Counter is the packet's AEAD message counter within that tunnel.
|
|
||||||
Counter uint64
|
|
||||||
}
|
|
||||||
|
|
||||||
// flowKey identifies a transport flow by {src, dst, sport, dport, family}.
|
|
||||||
// Comparable, so map lookups and linear scans over the slot list stay tight.
|
|
||||||
// Shared by the TCP and UDP coalescers; each coalescer keeps its own
|
|
||||||
// openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto never alias.
|
|
||||||
type flowKey struct {
|
|
||||||
src, dst [16]byte
|
|
||||||
sport, dport uint16
|
|
||||||
isV6 bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// initialSlots is the starting capacity of the slot pool.
|
|
||||||
// One flow per packet is the worst case, so this matches a typical carrier-side recvmmsg batch on the UDP socket.
|
|
||||||
const initialSlots = 64
|
|
||||||
|
|
||||||
// parseIPAt validates the IP header for lane parsing. newPacket already resolved the L4 protocol
|
|
||||||
// and offset for the firewall, so there is no proto sniff here; the caller's ipHdrLen is
|
|
||||||
// cross-checked instead. A plain header (v4 IHL 20, v6 exactly 40) is the only coalesceable
|
|
||||||
// shape. The v6 check is load-bearing: it rejects extension-header packets whose L4 is not at byte 40.
|
|
||||||
//
|
|
||||||
// The prologues fill fk's addresses and family in place (ports belong to the L4 parser; fk must
|
|
||||||
// be zero on entry so the v4 path leaves src[4:]/dst[4:] clear for map equality) and return pkt
|
|
||||||
// trimmed to the IP-declared length. The receiver-as-out-pointer shape is deliberate: these
|
|
||||||
// functions are too big to inline, and returning structs by value put five 64-byte copies on the
|
|
||||||
// per-packet path.
|
|
||||||
func (fk *flowKey) parseIPAt(pkt []byte, ipHdrLen int) ([]byte, bool) {
|
|
||||||
if len(pkt) < 20 {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
switch pkt[0] >> 4 {
|
|
||||||
case 4:
|
|
||||||
if ipHdrLen != 20 {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
return fk.parseIPv4Prologue(pkt)
|
|
||||||
case 6:
|
|
||||||
if ipHdrLen != 40 || len(pkt) < 40 {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
return fk.parseIPv6Prologue(pkt)
|
|
||||||
}
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseIPv4Prologue is the shared IPv4 tail of the prologue entries; the caller has verified
|
|
||||||
// len(pkt) >= 20 and the version.
|
|
||||||
func (fk *flowKey) parseIPv4Prologue(pkt []byte) ([]byte, bool) {
|
|
||||||
ihl := int(pkt[0]&0x0f) * 4
|
|
||||||
if ihl != 20 {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
// Reject any fragmentation (MF or nonzero offset). The dispatcher already gated FragAny; kept
|
|
||||||
// as defense in depth, since a fragment folded into a superpacket would corrupt reassembly.
|
|
||||||
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
|
|
||||||
if totalLen > len(pkt) || totalLen < ihl {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
fk.isV6 = false
|
|
||||||
copy(fk.src[:4], pkt[12:16])
|
|
||||||
copy(fk.dst[:4], pkt[16:20])
|
|
||||||
return pkt[:totalLen], true
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseIPv6Prologue is the shared IPv6 tail; the caller has verified len(pkt) >= 40, the version,
|
|
||||||
// and that the L4 header sits at byte 40.
|
|
||||||
func (fk *flowKey) parseIPv6Prologue(pkt []byte) ([]byte, bool) {
|
|
||||||
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
|
|
||||||
if 40+payloadLen > len(pkt) {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
fk.isV6 = true
|
|
||||||
copy(fk.src[:], pkt[8:24])
|
|
||||||
copy(fk.dst[:], pkt[24:40])
|
|
||||||
return pkt[:40+payloadLen], true
|
|
||||||
}
|
|
||||||
|
|
||||||
// ipHeadersMatch compares the IP portion of two packet header prefixes for
|
|
||||||
// byte-for-byte equality on every field that must be identical across coalesced segments.
|
|
||||||
// Size/IPID/IPCsum are masked out.
|
|
||||||
// The full DSCP/ECN byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel GRO:
|
|
||||||
// segments with differing ECN codepoints must not coalesce,
|
|
||||||
// otherwise ORing e.g. ECT(0) with ECT(1) would fabricate a false CE (congestion) mark or mark a Not-ECT flow as ECN-capable.
|
|
||||||
//
|
|
||||||
// The transport (L4) portion of the header is checked separately by the per-protocol matcher.
|
|
||||||
func ipHeadersMatch(a, b []byte, isV6 bool) bool {
|
|
||||||
if isV6 {
|
|
||||||
// IPv6: [0:4] = version/TC/flow label (TC[1:0] is ECN, so the full TC byte must match),
|
|
||||||
// [6:40] = next_hdr/hop + src + dst. Skip [4:6] payload_len.
|
|
||||||
return bytes.Equal(a[:4], b[:4]) && bytes.Equal(a[6:40], b[6:40])
|
|
||||||
}
|
|
||||||
// IPv4: [0:2] = version/IHL + DSCP|ECN (full ECN byte must match),
|
|
||||||
// [6:10] = flags/fragoff/TTL/proto, [12:20] = src+dst.
|
|
||||||
// Skip [2:4] total len, [4:6] id, [10:12] csum.
|
|
||||||
return bytes.Equal(a[:2], b[:2]) && bytes.Equal(a[6:10], b[6:10]) && bytes.Equal(a[12:20], b[12:20])
|
|
||||||
}
|
|
||||||
|
|
||||||
// ipv4FlagDF is the Don't Fragment bit in the IPv4 flags byte (header byte 6).
|
|
||||||
const ipv4FlagDF = 0x40
|
|
||||||
|
|
||||||
// ipv4CanCoalesceID reports whether an IPv4 packet whose header starts at
|
|
||||||
// nextHdr may join a chain whose seed header is seedHdr as segment index seg
|
|
||||||
// (the seed is segment 0). Kernel GSO re-stamps outgoing segment IDs as
|
|
||||||
// seed_id+n, so coalescing is only transparent when that re-stamp is either
|
|
||||||
// harmless (DF set: RFC 6864 atomic datagrams, the ID carries no meaning) or
|
|
||||||
// reproduces the original IDs exactly (DF clear + IDs already sequential —
|
|
||||||
// the same admission rule kernel GRO applies). Without this, a DF=0 sender
|
|
||||||
// with non-sequential IDs (e.g. OpenBSD's randomized IDs) could have IDs
|
|
||||||
// rewritten into ranges that collide across superpackets, corrupting
|
|
||||||
// reassembly if the packets are fragmented after the TUN write.
|
|
||||||
//
|
|
||||||
// DF itself is guaranteed uniform across a chain by ipHeadersMatch (byte 6
|
|
||||||
// is inside its compared range), so checking the seed's copy suffices.
|
|
||||||
func ipv4CanCoalesceID(seedHdr, nextHdr []byte, seg int) bool {
|
|
||||||
if seedHdr[6]&ipv4FlagDF != 0 {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
expect := binary.BigEndian.Uint16(seedHdr[4:6]) + uint16(seg)
|
|
||||||
return binary.BigEndian.Uint16(nextHdr[4:6]) == expect
|
|
||||||
}
|
|
||||||
|
|
||||||
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
|
|
||||||
// slices via Reserve and releases them in bulk via Reset.
|
|
||||||
type Arena struct {
|
|
||||||
buf []byte
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewArena returns an Arena with a pre-allocated backing of the given capacity.
|
|
||||||
func NewArena(capacity int) *Arena {
|
|
||||||
return &Arena{buf: make([]byte, 0, capacity)}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Reserve hands out a non-overlapping sz-byte slice from the arena.
|
|
||||||
// If the request doesn't fit the current backing, a fresh, larger backing is allocated.
|
|
||||||
// Already-borrowed slices reference the old backing and remain valid until Reset.
|
|
||||||
func (a *Arena) Reserve(sz int) []byte {
|
|
||||||
if len(a.buf)+sz > cap(a.buf) {
|
|
||||||
newCap := max(cap(a.buf)*2, sz)
|
|
||||||
a.buf = make([]byte, 0, newCap)
|
|
||||||
}
|
|
||||||
start := len(a.buf)
|
|
||||||
a.buf = a.buf[:start+sz]
|
|
||||||
return a.buf[start : start+sz : start+sz]
|
|
||||||
}
|
|
||||||
|
|
||||||
// Reset releases every slice handed out since the last Reset.
|
|
||||||
// Callers must not use any previously-borrowed slice after this returns.
|
|
||||||
// The underlying backing array is retained so subsequent Reserves don't re-allocate.
|
|
||||||
func (a *Arena) Reset() {
|
|
||||||
a.buf = a.buf[:0]
|
|
||||||
}
|
|
||||||
@@ -1,112 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
)
|
|
||||||
|
|
||||||
// stagePackets builds the stagedPacket entries Commit would have produced, so dispatch benchmarks
|
|
||||||
// bypass staging and the sort entirely.
|
|
||||||
func stagePackets(pkts [][]byte) []stagedPacket {
|
|
||||||
staged := make([]stagedPacket, len(pkts))
|
|
||||||
for i, p := range pkts {
|
|
||||||
pp := testPP(p)
|
|
||||||
staged[i] = stagedPacket{
|
|
||||||
pkt: p,
|
|
||||||
key: SortKey{Epoch: 1, Counter: uint64(i + 1)},
|
|
||||||
proto: pp.Protocol,
|
|
||||||
fragAny: pp.FragAny,
|
|
||||||
ipHdrLen: uint16(pp.IPHdrLen),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return staged
|
|
||||||
}
|
|
||||||
|
|
||||||
func flushLanes(b *testing.B, m *MultiCoalescer) {
|
|
||||||
b.Helper()
|
|
||||||
if m.tcp != nil {
|
|
||||||
if err := m.tcp.Flush(); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if m.udp != nil {
|
|
||||||
if err := m.udp.Flush(); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := m.pt.Flush(); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// runDispatchBench measures dispatch plus the per-batch lane flush: the post-sort half of the
|
|
||||||
// batcher, which is where the production profile concentrates.
|
|
||||||
func runDispatchBench(b *testing.B, pkts [][]byte, batchSize int) {
|
|
||||||
b.Helper()
|
|
||||||
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
|
|
||||||
staged := stagePackets(pkts)
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.SetBytes(int64(len(pkts[0])))
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
if err := m.dispatch(staged[i%len(staged)]); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
if (i+1)%batchSize == 0 {
|
|
||||||
flushLanes(b, m)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
b.StopTimer()
|
|
||||||
flushLanes(b, m)
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkDispatchSingleFlow is the bulk steady state: every packet past the seed appends.
|
|
||||||
func BenchmarkDispatchSingleFlow(b *testing.B) {
|
|
||||||
runDispatchBench(b, buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200), tcpCoalesceMaxSegs)
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkDispatchInterleaved16 stresses the openSlots map: 16 flows round-robined defeats the
|
|
||||||
// lastSlot cache on every packet.
|
|
||||||
func BenchmarkDispatchInterleaved16(b *testing.B) {
|
|
||||||
pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200)
|
|
||||||
runDispatchBench(b, pkts, len(pkts))
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkDispatchAckHeavy alternates MSS data with pure ACKs on one flow — the RX shape of a
|
|
||||||
// bidirectional transfer (the peer's data and its ACKs of our data share the tunnel direction).
|
|
||||||
func BenchmarkDispatchAckHeavy(b *testing.B) {
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
var pkts [][]byte
|
|
||||||
seq := uint32(1000)
|
|
||||||
for range tcpCoalesceMaxSegs / 2 {
|
|
||||||
pkts = append(pkts, buildTCPv4(seq, tcpAck, pay))
|
|
||||||
seq += uint32(len(pay))
|
|
||||||
pkts = append(pkts, buildTCPv4(seq, tcpAck, nil))
|
|
||||||
}
|
|
||||||
runDispatchBench(b, pkts, len(pkts))
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkDispatchUDPFlow is the QUIC-ish bulk UDP shape.
|
|
||||||
func BenchmarkDispatchUDPFlow(b *testing.B) {
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
pkts := make([][]byte, udpCoalesceMaxSegs)
|
|
||||||
for i := range pkts {
|
|
||||||
pkts[i] = buildUDPv4(2000, 443, pay)
|
|
||||||
}
|
|
||||||
runDispatchBench(b, pkts, len(pkts))
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkDispatchSeedHeavy sets PSH on every packet so each one seeds and immediately closes
|
|
||||||
// its own slot — the small-write RPC shape, and the upper bound on what the seed path (including
|
|
||||||
// the parsedTCP-to-slot field transfer) can cost.
|
|
||||||
func BenchmarkDispatchSeedHeavy(b *testing.B) {
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
pkts := make([][]byte, tcpCoalesceMaxSegs)
|
|
||||||
seq := uint32(1000)
|
|
||||||
for i := range pkts {
|
|
||||||
pkts[i] = buildTCPv4(seq, tcpAckPsh, pay)
|
|
||||||
seq += uint32(len(pay))
|
|
||||||
}
|
|
||||||
runDispatchBench(b, pkts, len(pkts))
|
|
||||||
}
|
|
||||||
@@ -1,76 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
//TODO refactor this away
|
|
||||||
// This file holds the lanes' self-parsing Commit entries and the proto-checking parsers behind
|
|
||||||
// them. Production traffic enters the lanes only through MultiCoalescer.dispatch and the At
|
|
||||||
// parsers; these wrappers reproduce that path (including seal-all on unparseable shapes) on top
|
|
||||||
// of a local parse, so tests and benches can drive one lane with nothing but a packet.
|
|
||||||
|
|
||||||
// parseIPPrologue resolves the IP version, requires the L4 protocol to match wantProto (6 TCP,
|
|
||||||
// 17 UDP), and defers to the shared per-version cores. Returns the trimmed packet and the L4
|
|
||||||
// offset; fk must be zero on entry and is filled in place.
|
|
||||||
func (fk *flowKey) parseIPPrologue(pkt []byte, wantProto byte) ([]byte, int, bool) {
|
|
||||||
if len(pkt) < 20 {
|
|
||||||
return nil, 0, false
|
|
||||||
}
|
|
||||||
switch pkt[0] >> 4 {
|
|
||||||
case 4:
|
|
||||||
if pkt[9] != wantProto {
|
|
||||||
return nil, 0, false
|
|
||||||
}
|
|
||||||
trimmed, ok := fk.parseIPv4Prologue(pkt)
|
|
||||||
return trimmed, 20, ok
|
|
||||||
case 6:
|
|
||||||
if len(pkt) < 40 {
|
|
||||||
return nil, 0, false
|
|
||||||
}
|
|
||||||
if pkt[6] != wantProto {
|
|
||||||
return nil, 0, false
|
|
||||||
}
|
|
||||||
trimmed, ok := fk.parseIPv6Prologue(pkt)
|
|
||||||
return trimmed, 40, ok
|
|
||||||
}
|
|
||||||
return nil, 0, false
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseBase extracts the flow key and IP/TCP offsets for any TCP packet, admissible for
|
|
||||||
// coalescing or not. Returns false for non-TCP or malformed input.
|
|
||||||
func (p *parsedTCP) parseBase(pkt []byte) bool {
|
|
||||||
trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoTCP)
|
|
||||||
if !ok {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return p.parseTail(trimmed, ipHdrLen)
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseBase extracts the flow key and IP/UDP offsets for a UDP packet.
|
|
||||||
func (p *parsedUDP) parseBase(pkt []byte) bool {
|
|
||||||
trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoUDP)
|
|
||||||
if !ok {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return p.parseTail(trimmed, ipHdrLen)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
|
||||||
func (c *TCPCoalescer) Commit(pkt []byte) error {
|
|
||||||
var info parsedTCP
|
|
||||||
if !info.parseBase(pkt) {
|
|
||||||
// Unparseable: flow key unknown, seal everything so later data cannot emit ahead of it.
|
|
||||||
c.sealAllOpen()
|
|
||||||
c.addVerbatim(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return c.commitParsed(pkt, &info)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
|
||||||
func (c *UDPCoalescer) Commit(pkt []byte) error {
|
|
||||||
var info parsedUDP
|
|
||||||
if !info.parseBase(pkt) {
|
|
||||||
c.sealAllOpen()
|
|
||||||
c.addVerbatim(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return c.commitParsed(pkt, &info)
|
|
||||||
}
|
|
||||||
@@ -1,133 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"cmp"
|
|
||||||
"errors"
|
|
||||||
"io"
|
|
||||||
"log/slog"
|
|
||||||
"slices"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
|
||||||
)
|
|
||||||
|
|
||||||
// MultiCoalescer stages plaintext packets with their (epoch, counter) sort keys and, at Flush,
|
|
||||||
// replays them in sender-transmission order into lane-specific batchers selected by L4 protocol.
|
|
||||||
//
|
|
||||||
// Sorting before dispatch keeps the ordering story simple: each lane consumes packets in
|
|
||||||
// transmission order, builds slots in that order, and emits them in creation order. Wire reorder
|
|
||||||
// inside a flush batch is repaired here, before it can fragment a lane's coalesce chains, so the
|
|
||||||
// lanes carry no reorder-repair machinery.
|
|
||||||
//
|
|
||||||
// The contract is per-tunnel transmission order within each lane, with two exceptions: a pure TCP
|
|
||||||
// ACK may be overtaken by later same-flow data (it does not close the flow's open chain; a late
|
|
||||||
// ACK is just a stale ACK), and an unparseable shape seals every open chain in its lane (its flow
|
|
||||||
// is unknown) and rides the lane as an in-lane verbatim, still in transmission order. Routing
|
|
||||||
// follows the flow: a flow's non-coalesceable shapes ride its protocol lane rather than falling
|
|
||||||
// to the later-flushed pt lane.
|
|
||||||
//
|
|
||||||
// Cross-lane order (TCP vs UDP vs everything else) is not preserved.
|
|
||||||
type MultiCoalescer struct {
|
|
||||||
tcp *TCPCoalescer
|
|
||||||
udp *UDPCoalescer
|
|
||||||
pt *Passthrough
|
|
||||||
|
|
||||||
// staged holds this batch's packets and sort keys until Flush. Borrowed: the caller keeps
|
|
||||||
// each pkt alive until Flush returns.
|
|
||||||
staged []stagedPacket
|
|
||||||
}
|
|
||||||
|
|
||||||
// stagedPacket carries the scalars dispatch needs from the firewall's ParsedPacket, copied by
|
|
||||||
// value: pp is reused by the caller per packet and must not be retained past Commit.
|
|
||||||
type stagedPacket struct {
|
|
||||||
pkt []byte
|
|
||||||
key SortKey
|
|
||||||
proto byte
|
|
||||||
fragAny bool
|
|
||||||
ipHdrLen uint16
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewMultiCoalescer builds a multi-lane batcher over w, based on available protocol support. The
|
|
||||||
// staging sort applies even when no GSO lane is available: passthrough-only platforms still get
|
|
||||||
// transmission-order repair.
|
|
||||||
func NewMultiCoalescer(w io.Writer, l *slog.Logger) *MultiCoalescer {
|
|
||||||
m := &MultiCoalescer{
|
|
||||||
pt: NewPassthrough(w),
|
|
||||||
staged: make([]stagedPacket, 0, initialSlots),
|
|
||||||
}
|
|
||||||
m.tcp = NewTCPCoalescer(w, l)
|
|
||||||
m.udp = NewUDPCoalescer(w)
|
|
||||||
return m
|
|
||||||
}
|
|
||||||
|
|
||||||
// Commit stages pkt for the next Flush; dispatch is deferred so it runs on packets already in
|
|
||||||
// transmission order. key carries the packet's tunnel epoch and message counter. pkt is borrowed:
|
|
||||||
// the caller must keep it valid until the next Flush and not re-use it, and Flush may patch a
|
|
||||||
// coalesced packet's headers in place. pp is the firewall's parse of pkt and is borrowed only
|
|
||||||
// for this call, so the fields dispatch needs are copied here.
|
|
||||||
func (m *MultiCoalescer) Commit(pkt []byte, key SortKey, pp *firewall.ParsedPacket) error {
|
|
||||||
m.staged = append(m.staged, stagedPacket{
|
|
||||||
pkt: pkt,
|
|
||||||
key: key,
|
|
||||||
proto: pp.Protocol,
|
|
||||||
fragAny: pp.FragAny,
|
|
||||||
ipHdrLen: uint16(pp.IPHdrLen),
|
|
||||||
})
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// compareStaged orders staged packets by (epoch, counter)
|
|
||||||
func compareStaged(a, b stagedPacket) int {
|
|
||||||
if c := cmp.Compare(a.key.Epoch, b.key.Epoch); c != 0 {
|
|
||||||
return c
|
|
||||||
}
|
|
||||||
return cmp.Compare(a.key.Counter, b.key.Counter)
|
|
||||||
}
|
|
||||||
|
|
||||||
// dispatch routes one staged packet to its protocol lane (see commitStaged), or to the verbatim
|
|
||||||
// passthrough when the lane has no GSO support.
|
|
||||||
func (m *MultiCoalescer) dispatch(sp stagedPacket) error {
|
|
||||||
switch sp.proto {
|
|
||||||
case ipProtoTCP:
|
|
||||||
if m.tcp != nil {
|
|
||||||
return m.tcp.commitStaged(sp)
|
|
||||||
}
|
|
||||||
case ipProtoUDP:
|
|
||||||
if m.udp != nil {
|
|
||||||
return m.udp.commitStaged(sp)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return m.pt.enqueue(sp.pkt)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Flush sorts the staged batch into transmission order, replays it into the lanes, then flushes each lane.
|
|
||||||
// Drains everything and returns the joined errors; one bad packet does not hold up the rest.
|
|
||||||
// After Flush returns, committed payload slices may be recycled.
|
|
||||||
func (m *MultiCoalescer) Flush() error {
|
|
||||||
// Arrival order is already almost sorted (reorder is the exception), which pdqsort detects
|
|
||||||
// and handles in near-linear time.
|
|
||||||
slices.SortFunc(m.staged, compareStaged)
|
|
||||||
|
|
||||||
var errs []error
|
|
||||||
for _, sp := range m.staged {
|
|
||||||
if err := m.dispatch(sp); err != nil {
|
|
||||||
errs = append(errs, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
clear(m.staged) // drop borrowed pkt refs
|
|
||||||
m.staged = m.staged[:0]
|
|
||||||
|
|
||||||
if m.tcp != nil {
|
|
||||||
if err := m.tcp.Flush(); err != nil {
|
|
||||||
errs = append(errs, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if m.udp != nil {
|
|
||||||
if err := m.udp.Flush(); err != nil {
|
|
||||||
errs = append(errs, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := m.pt.Flush(); err != nil {
|
|
||||||
errs = append(errs, err)
|
|
||||||
}
|
|
||||||
return errors.Join(errs...)
|
|
||||||
}
|
|
||||||
@@ -1,437 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"encoding/binary"
|
|
||||||
"io"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
)
|
|
||||||
|
|
||||||
// keySeq hands out SortKeys with ascending counters in a fixed epoch, for
|
|
||||||
// tests where commit order IS transmission order.
|
|
||||||
type keySeq struct {
|
|
||||||
epoch, counter uint64
|
|
||||||
}
|
|
||||||
|
|
||||||
func (k *keySeq) next() SortKey {
|
|
||||||
k.counter++
|
|
||||||
return SortKey{Epoch: k.epoch, Counter: k.counter}
|
|
||||||
}
|
|
||||||
|
|
||||||
// newTestMultiCoalescer builds a batcher over w.
|
|
||||||
func newTestMultiCoalescer(tb testing.TB, w io.Writer) *MultiCoalescer {
|
|
||||||
tb.Helper()
|
|
||||||
return NewMultiCoalescer(w, test.NewLogger())
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right
|
|
||||||
// lane: TCP and UDP get coalesced when their lanes are enabled, anything
|
|
||||||
// else (ICMP here) falls through to plain Write.
|
|
||||||
func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
m := newTestMultiCoalescer(t, w)
|
|
||||||
k := &keySeq{epoch: 1}
|
|
||||||
|
|
||||||
tcpPay := make([]byte, 1200)
|
|
||||||
udpPay := make([]byte, 1200)
|
|
||||||
icmp := make([]byte, 28)
|
|
||||||
icmp[0] = 0x45
|
|
||||||
icmp[2] = 0
|
|
||||||
icmp[3] = 28
|
|
||||||
icmp[9] = 1
|
|
||||||
|
|
||||||
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(1000, tcpAck, tcpPay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(2200, tcpAck, tcpPay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(icmp, k.next(), testPP(icmp)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// 1 TCP super (2 segments) + 1 UDP super (2 segments) = 2 gso writes.
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes (one TCP + one UDP), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 {
|
|
||||||
t.Fatalf("want 1 plain write (ICMP), got %d", len(w.writes))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMultiCoalescerRestoresTransmissionOrder is the core staging-sort
|
|
||||||
// property: packets committed out of counter order (wire reorder inside one
|
|
||||||
// flush batch) are replayed into the lanes in transmission order, so the
|
|
||||||
// reorder never fragments the coalesce chain — one superpacket, in seq
|
|
||||||
// order, exactly as if the wire had never reordered. The retransmit shape
|
|
||||||
// falls out of the same key: a retransmit carries a lower seq but a HIGHER
|
|
||||||
// counter (it was encrypted later), so it emits after the data it trails.
|
|
||||||
func TestMultiCoalescerRestoresTransmissionOrder(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
m := newTestMultiCoalescer(t, w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
|
|
||||||
// Transmission order: seq 1000 (c1), 2200 (c2), 3400 (c3).
|
|
||||||
// Arrival order: 3400, 1000, 2200.
|
|
||||||
if err := m.Commit(buildTCPv4(3400, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(3400, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
|
||||||
t.Fatalf("want 1 gso write (unfragmented chain), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
if len(g.pays) != 3 {
|
|
||||||
t.Fatalf("segs=%d want 3", len(g.pays))
|
|
||||||
}
|
|
||||||
const ipHdrLen = 20
|
|
||||||
if seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 {
|
|
||||||
t.Errorf("seed seq=%d want 1000", seedSeq)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Retransmit: seq 1000 again but counter 4 — sorts after seq 4600 (c3).
|
|
||||||
w.writes, w.gsoWrites, w.order = nil, nil, nil
|
|
||||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildTCPv4(4600, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(4600, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 2 {
|
|
||||||
t.Fatalf("want 2 plain writes, got %d (gso=%d)", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
first := binary.BigEndian.Uint32(w.writes[0][24:28])
|
|
||||||
second := binary.BigEndian.Uint32(w.writes[1][24:28])
|
|
||||||
if first != 4600 || second != 1000 {
|
|
||||||
t.Fatalf("emission (%d, %d), want (4600, 1000): retransmit must not overtake in-flight data", first, second)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMultiCoalescerRestoresOrderAcrossFlows scrambles two interleaved flows;
|
|
||||||
// the staging sort must repair each flow into one superpacket without any
|
|
||||||
// cross-flow contamination.
|
|
||||||
func TestMultiCoalescerRestoresOrderAcrossFlows(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
m := newTestMultiCoalescer(t, w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
|
|
||||||
// Transmission: A.100 (c1), B.500 (c2), A.1300 (c3), B.1700 (c4).
|
|
||||||
// Arrival: A.1300, B.1700, A.100, B.500.
|
|
||||||
if err := m.Commit(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes (one per flow), got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
|
||||||
}
|
|
||||||
for i, g := range w.gsoWrites {
|
|
||||||
if len(g.pays) != 2 {
|
|
||||||
t.Errorf("gso[%d] segs=%d want 2", i, len(g.pays))
|
|
||||||
}
|
|
||||||
const ipHdrLen = 20
|
|
||||||
seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8])
|
|
||||||
sport := binary.BigEndian.Uint16(g.hdr[ipHdrLen : ipHdrLen+2])
|
|
||||||
switch sport {
|
|
||||||
case 1000:
|
|
||||||
if seedSeq != 100 {
|
|
||||||
t.Errorf("flow A seed seq=%d want 100", seedSeq)
|
|
||||||
}
|
|
||||||
case 3000:
|
|
||||||
if seedSeq != 500 {
|
|
||||||
t.Errorf("flow B seed seq=%d want 500", seedSeq)
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
t.Errorf("unexpected sport %d", sport)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMultiCoalescerEpochOrdersAcrossRehandshake: a re-handshake replaces
|
|
||||||
// the tunnel, and the replacement's counter space starts near zero — raw
|
|
||||||
// counter order would emit the new tunnel's packets first while the old
|
|
||||||
// tunnel's backlog is still arriving. The epoch key must dominate:
|
|
||||||
// everything from the old tunnel emits before anything from the new one.
|
|
||||||
func TestMultiCoalescerEpochOrdersAcrossRehandshake(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
m := newTestMultiCoalescer(t, w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
|
|
||||||
// New session's first data arrives before the old session's last data.
|
|
||||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 8, Counter: 1}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 7, Counter: 9_000_000}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Same flow, contiguous seq, identical headers: after the epoch sort the
|
|
||||||
// two segments append into one superpacket seeded by the OLD session's
|
|
||||||
// packet.
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
|
||||||
}
|
|
||||||
const ipHdrLen = 20
|
|
||||||
if seedSeq := binary.BigEndian.Uint32(w.gsoWrites[0].hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 {
|
|
||||||
t.Errorf("seed seq=%d want 1000 (old session first)", seedSeq)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMultiCoalescerNoUSOFallsThrough verifies that on a queue without USO
|
|
||||||
// (older kernel: TSO but no GSO_UDP_L4) the UDP lane never comes up and UDP
|
|
||||||
// packets still reach the kernel via verbatim rather than being lost.
|
|
||||||
func TestMultiCoalescerNoUSOFallsThrough(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true, noUSO: true}
|
|
||||||
m := newTestMultiCoalescer(t, w)
|
|
||||||
k := &keySeq{epoch: 1}
|
|
||||||
if m.udp != nil {
|
|
||||||
t.Fatal("UDP lane must not come up without USO")
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 0 {
|
|
||||||
t.Errorf("UDP must NOT be coalesced when USO disabled, got %d gso writes", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.writes) != 2 {
|
|
||||||
t.Errorf("UDP must pass through as 2 plain writes, got %d", len(w.writes))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMultiCoalescerNoOffloadsStillSorts covers a queue that can't offload
|
|
||||||
// anything. Both lane constructors refuse, so every packet rides the
|
|
||||||
// verbatim lane — but the staging sort still applies, so emission follows
|
|
||||||
// transmission order even without GSO.
|
|
||||||
func TestMultiCoalescerNoOffloadsStillSorts(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: false}
|
|
||||||
m := newTestMultiCoalescer(t, w)
|
|
||||||
if m.tcp != nil || m.udp != nil {
|
|
||||||
t.Fatal("no lane may come up without offloads")
|
|
||||||
}
|
|
||||||
pkts := [][]byte{
|
|
||||||
buildTCPv4(1000, tcpAck, make([]byte, 1200)),
|
|
||||||
buildUDPv4(1000, 53, make([]byte, 800)),
|
|
||||||
buildTCPv4(2200, tcpAck, make([]byte, 1200)),
|
|
||||||
}
|
|
||||||
// Committed in reverse transmission order; keys carry the truth.
|
|
||||||
for i := len(pkts) - 1; i >= 0; i-- {
|
|
||||||
if err := m.Commit(pkts[i], SortKey{Epoch: 1, Counter: uint64(i + 1)}, testPP(pkts[i])); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := m.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 0 {
|
|
||||||
t.Errorf("no GSO writes possible, got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.writes) != len(pkts) {
|
|
||||||
t.Fatalf("want %d plain writes, got %d", len(pkts), len(w.writes))
|
|
||||||
}
|
|
||||||
// One lane for everything means the sorted order survives end to end.
|
|
||||||
for i, want := range pkts {
|
|
||||||
if !bytes.Equal(w.writes[i], want) {
|
|
||||||
t.Errorf("write %d out of order or corrupt", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildUDPv6Fragment builds an IPv6 packet whose extension chain is a
|
|
||||||
// single fragment header (NH=44) naming UDP as the terminal protocol —
|
|
||||||
// a first fragment (offset 0, MF set) carrying the UDP header and a
|
|
||||||
// partial payload.
|
|
||||||
func buildUDPv6Fragment(sport, dport uint16, payload []byte) []byte {
|
|
||||||
const ipHdrLen = 40
|
|
||||||
const fragHdrLen = 8
|
|
||||||
const udpHdrLen = 8
|
|
||||||
total := ipHdrLen + fragHdrLen + udpHdrLen + len(payload)
|
|
||||||
pkt := make([]byte, total)
|
|
||||||
|
|
||||||
pkt[0] = 0x60
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(total-ipHdrLen))
|
|
||||||
pkt[6] = 44 // fragment extension header
|
|
||||||
pkt[7] = 64
|
|
||||||
pkt[8] = 0xfe
|
|
||||||
pkt[9] = 0x80
|
|
||||||
pkt[23] = 1
|
|
||||||
pkt[24] = 0xfe
|
|
||||||
pkt[25] = 0x80
|
|
||||||
pkt[39] = 2
|
|
||||||
|
|
||||||
pkt[40] = ipProtoUDP // fragment's next header
|
|
||||||
binary.BigEndian.PutUint16(pkt[42:44], 0x0001) // offset 0, MF set
|
|
||||||
binary.BigEndian.PutUint32(pkt[44:48], 0x1badf00) // identification
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(pkt[48:50], sport)
|
|
||||||
binary.BigEndian.PutUint16(pkt[50:52], dport)
|
|
||||||
binary.BigEndian.PutUint16(pkt[52:54], uint16(udpHdrLen+len(payload)))
|
|
||||||
copy(pkt[56:], payload)
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMultiCoalescerIPv6FragmentStaysInLane locks in extension-header
|
|
||||||
// routing: a fragment whose chain terminates in UDP must ride the UDP lane
|
|
||||||
// as an in-lane verbatim — emitted ahead of later same-flow datagrams —
|
|
||||||
// not the verbatim lane, which flushes after every coalescer lane and
|
|
||||||
// would reorder it behind data that arrived after it.
|
|
||||||
func TestMultiCoalescerIPv6FragmentStaysInLane(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
m := newTestMultiCoalescer(t, w)
|
|
||||||
k := &keySeq{epoch: 1}
|
|
||||||
|
|
||||||
if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 {
|
|
||||||
t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes))
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want the two whole datagrams coalesced into 1 gso write, got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
// Transmission order was fragment-then-data; same-lane routing must keep it.
|
|
||||||
if w.order[0] != "write" {
|
|
||||||
t.Fatalf("fragment must be emitted before later data (in-lane verbatim), order=%v", w.order)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMultiCoalescerFragmentSealsUDPChains: an unparseable datagram
|
|
||||||
// (fragment) seals every open UDP chain, so datagrams from before and after
|
|
||||||
// it land in separate superpackets and the fragment holds its transmission-
|
|
||||||
// order position between them.
|
|
||||||
func TestMultiCoalescerFragmentSealsUDPChains(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
m := newTestMultiCoalescer(t, w)
|
|
||||||
k := &keySeq{epoch: 1}
|
|
||||||
|
|
||||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes (chains sealed around the fragment), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 {
|
|
||||||
t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes))
|
|
||||||
}
|
|
||||||
want := []string{"gso", "write", "gso"}
|
|
||||||
if len(w.order) != 3 || w.order[0] != want[0] || w.order[1] != want[1] || w.order[2] != want[2] {
|
|
||||||
t.Fatalf("emission order = %v, want %v", w.order, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMultiCoalescerNoTSOFallsThrough mirrors the no-TSO case.
|
|
||||||
func TestMultiCoalescerNoTSOFallsThrough(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true, noTSO: true}
|
|
||||||
m := newTestMultiCoalescer(t, w)
|
|
||||||
k := &keySeq{epoch: 1}
|
|
||||||
if m.tcp != nil {
|
|
||||||
t.Fatal("TCP lane must not come up without TSO")
|
|
||||||
}
|
|
||||||
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), k.next(), testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), k.next(), testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := m.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 0 {
|
|
||||||
t.Errorf("TCP must NOT be coalesced when TSO disabled, got %d gso writes", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.writes) != 2 {
|
|
||||||
t.Errorf("TCP must pass through as 2 plain writes, got %d", len(w.writes))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// testPP derives the ParsedPacket newPacket would produce for the packet
|
|
||||||
// shapes the tests build: plain v4/v6, v4 with options or fragment bits set,
|
|
||||||
// and the single-fragment-header v6 shape from buildUDPv6Fragment. Anything
|
|
||||||
// unrecognizable stays zero (proto 0 routes to the passthrough lane).
|
|
||||||
func testPP(pkt []byte) *firewall.ParsedPacket {
|
|
||||||
pp := &firewall.ParsedPacket{}
|
|
||||||
if len(pkt) < 20 {
|
|
||||||
return pp
|
|
||||||
}
|
|
||||||
switch pkt[0] >> 4 {
|
|
||||||
case 4:
|
|
||||||
pp.Protocol = pkt[9]
|
|
||||||
pp.IPHdrLen = int(pkt[0]&0x0f) * 4
|
|
||||||
pp.FragAny = binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0
|
|
||||||
case 6:
|
|
||||||
pp.Protocol = pkt[6]
|
|
||||||
pp.IPHdrLen = 40
|
|
||||||
if pp.Protocol == 44 { // fragment extension header
|
|
||||||
pp.Protocol = pkt[40]
|
|
||||||
pp.IPHdrLen = 48
|
|
||||||
pp.FragAny = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return pp
|
|
||||||
}
|
|
||||||
@@ -1,38 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"io"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Passthrough is MultiCoalescer's verbatim lane: no batching, packets are written at Flush in the
|
|
||||||
// order enqueued.
|
|
||||||
type Passthrough struct {
|
|
||||||
out io.Writer
|
|
||||||
slots [][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewPassthrough(w io.Writer) *Passthrough {
|
|
||||||
return &Passthrough{
|
|
||||||
out: w,
|
|
||||||
slots: make([][]byte, 0, 128),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// enqueue accepts one packet, already sorted into transmission order by dispatch.
|
|
||||||
func (p *Passthrough) enqueue(pkt []byte) error {
|
|
||||||
p.slots = append(p.slots, pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (p *Passthrough) Flush() error {
|
|
||||||
var firstErr error
|
|
||||||
for _, s := range p.slots {
|
|
||||||
_, err := p.out.Write(s)
|
|
||||||
if err != nil && firstErr == nil {
|
|
||||||
firstErr = err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
clear(p.slots)
|
|
||||||
p.slots = p.slots[:0]
|
|
||||||
return firstErr
|
|
||||||
}
|
|
||||||
@@ -1,472 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"encoding/binary"
|
|
||||||
"io"
|
|
||||||
"log/slog"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
)
|
|
||||||
|
|
||||||
// ipProtoTCP is the IANA protocol number for TCP. Defined here to help Windows out.
|
|
||||||
const ipProtoTCP = 6
|
|
||||||
|
|
||||||
// tcpCoalesceBufSize caps total bytes per superpacket. Mirrors the kernel's
|
|
||||||
// sk_gso_max_size of ~64KiB; anything beyond this would be rejected anyway.
|
|
||||||
const tcpCoalesceBufSize = 65535
|
|
||||||
|
|
||||||
// tcpCoalesceMaxSegs caps how many segments we'll coalesce into a single
|
|
||||||
// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency.
|
|
||||||
const tcpCoalesceMaxSegs = 64
|
|
||||||
|
|
||||||
// coalesceSlot is one entry in the coalescer's ordered event queue. A verbatim slot holds a single
|
|
||||||
// borrowed packet emitted as-is (pure ACK, non-admissible TCP, unparseable, or oversize seed); a
|
|
||||||
// non-verbatim slot is an in-progress coalesced superpacket. payIovs are borrowed slices of the
|
|
||||||
// caller's plaintext buffers; the caller must keep them alive until Flush.
|
|
||||||
type coalesceSlot struct {
|
|
||||||
verbatim bool
|
|
||||||
// rawPkt is borrowed: the whole packet for verbatim slots, the seed packet for coalesce
|
|
||||||
// slots. A slot that never grows past one segment is emitted from rawPkt so its original
|
|
||||||
// (already valid) L4 checksum ships DATA_VALID instead of making the kernel recompute it.
|
|
||||||
// A multi-segment slot's superpacket header is rawPkt's, patched in place at flush.
|
|
||||||
rawPkt []byte
|
|
||||||
|
|
||||||
fk flowKey
|
|
||||||
hdrLen int
|
|
||||||
ipHdrLen int
|
|
||||||
isV6 bool
|
|
||||||
gsoSize int
|
|
||||||
numSeg int
|
|
||||||
totalPay int
|
|
||||||
nextSeq uint32
|
|
||||||
payIovs [][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
// TCPCoalescer accumulates adjacent in-flow TCP data segments across multiple concurrent flows and
|
|
||||||
// emits each flow's run as a single TSO superpacket via tio.GSOWriter. Input must be in sender
|
|
||||||
// transmission order (MultiCoalescer sorts by (epoch, counter) before dispatch); slots are emitted
|
|
||||||
// in creation order, so emission reproduces transmission order except for the pure-ACK case in
|
|
||||||
// commitParsed. Owns no locks; one coalescer per TUN write queue.
|
|
||||||
type TCPCoalescer struct {
|
|
||||||
w tio.GSOWriter
|
|
||||||
|
|
||||||
// slots is the ordered event queue. Flush walks it once and emits each
|
|
||||||
// entry as either a WriteGSO (coalesced) or a w.Write (verbatim).
|
|
||||||
slots []*coalesceSlot
|
|
||||||
// openSlots maps a flow key to its open slot so new segments can extend an in-progress
|
|
||||||
// superpacket in O(1). Removal is what closes a chain: on PSH or a short last segment, on a
|
|
||||||
// non-admissible packet for the flow, or in Flush.
|
|
||||||
openSlots map[flowKey]*coalesceSlot
|
|
||||||
// lastSlot caches the most recently touched open slot. Bulk traffic
|
|
||||||
// arrives in same-flow runs (single-flow steady state, or GRO bursts
|
|
||||||
// under multi-flow), so comparing the incoming key against the cached
|
|
||||||
// slot's own fk lets the hot path skip the map lookup (and the aeshash
|
|
||||||
// of a 38-byte key) for the length of each run.
|
|
||||||
// Kept in lockstep with openSlots: nil whenever the slot it pointed
|
|
||||||
// at is removed.
|
|
||||||
lastSlot *coalesceSlot
|
|
||||||
pool []*coalesceSlot // free list for reuse
|
|
||||||
l *slog.Logger
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewTCPCoalescer wraps w, returning nil if w can't accept GSO_TCP writes.
|
|
||||||
func NewTCPCoalescer(w io.Writer, l *slog.Logger) *TCPCoalescer {
|
|
||||||
gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP)
|
|
||||||
if !ok {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return &TCPCoalescer{
|
|
||||||
w: gw,
|
|
||||||
slots: make([]*coalesceSlot, 0, initialSlots),
|
|
||||||
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
|
|
||||||
pool: make([]*coalesceSlot, 0, initialSlots),
|
|
||||||
l: l,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// parsedTCP holds the fields extracted from a single parse so later steps
|
|
||||||
// (admission, slot lookup, canAppend) don't re-walk the header.
|
|
||||||
type parsedTCP struct {
|
|
||||||
fk flowKey
|
|
||||||
ipHdrLen int
|
|
||||||
hdrLen int
|
|
||||||
payLen int
|
|
||||||
seq uint32
|
|
||||||
flags byte
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseAt extracts the flow key and IP/TCP offsets for a packet the dispatcher already knows is
|
|
||||||
// TCP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
|
|
||||||
// entry and is filled in place; see flowKey.parseIPAt for why. Returns false for malformed input
|
|
||||||
// or any shape that must not coalesce (IPv4 options/fragmentation, IPv6 extension headers).
|
|
||||||
func (p *parsedTCP) parseAt(pkt []byte, ipHdrLen int) bool {
|
|
||||||
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
|
|
||||||
if !ok {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return p.parseTail(trimmed, ipHdrLen)
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseTail layers the TCP-header parse on a validated IP prologue. pkt is the trimmed packet;
|
|
||||||
// fk's addresses are already filled.
|
|
||||||
func (p *parsedTCP) parseTail(pkt []byte, ipHdrLen int) bool {
|
|
||||||
if len(pkt) < ipHdrLen+20 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
tcpOff := int(pkt[ipHdrLen+12]>>4) * 4
|
|
||||||
if tcpOff < 20 || tcpOff > 60 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if len(pkt) < ipHdrLen+tcpOff {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
p.ipHdrLen = ipHdrLen
|
|
||||||
p.hdrLen = ipHdrLen + tcpOff
|
|
||||||
p.payLen = len(pkt) - p.hdrLen
|
|
||||||
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
|
|
||||||
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
|
|
||||||
p.seq = binary.BigEndian.Uint32(pkt[ipHdrLen+4 : ipHdrLen+8])
|
|
||||||
p.flags = pkt[ipHdrLen+13]
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// TCP flag bits (byte 13 of the TCP header). Only the bits the coalescer consults are named;
|
|
||||||
// FIN/SYN/RST/URG/CWR are rejected by the negative mask in commitParsed.
|
|
||||||
const (
|
|
||||||
tcpFlagPsh = 0x08
|
|
||||||
tcpFlagAck = 0x10
|
|
||||||
tcpFlagEce = 0x40
|
|
||||||
)
|
|
||||||
|
|
||||||
// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
|
|
||||||
// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
|
|
||||||
func (c *TCPCoalescer) sealAllOpen() {
|
|
||||||
clear(c.openSlots)
|
|
||||||
c.lastSlot = nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
|
|
||||||
// hashing the 38-byte key when no chains are open (e.g. ack-dominant queues).
|
|
||||||
func (c *TCPCoalescer) sealFlow(fk flowKey) {
|
|
||||||
if len(c.openSlots) == 0 {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if last := c.lastSlot; last != nil && last.fk == fk {
|
|
||||||
c.lastSlot = nil
|
|
||||||
}
|
|
||||||
delete(c.openSlots, fk)
|
|
||||||
}
|
|
||||||
|
|
||||||
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
|
|
||||||
// coalesce (any fragmentation, unparseable header) seals every open chain
|
|
||||||
// and rides the lane as an in-lane verbatim, still in transmission order.
|
|
||||||
func (c *TCPCoalescer) commitStaged(sp stagedPacket) error {
|
|
||||||
if sp.fragAny {
|
|
||||||
c.sealAllOpen()
|
|
||||||
c.addVerbatim(sp.pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
var info parsedTCP
|
|
||||||
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
|
||||||
c.sealAllOpen()
|
|
||||||
c.addVerbatim(sp.pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return c.commitParsed(sp.pkt, &info)
|
|
||||||
}
|
|
||||||
|
|
||||||
// commitParsed commits one parsed TCP packet. The caller (dispatch, via parseAt) supplies a
|
|
||||||
// valid parse so the header is not re-walked here.
|
|
||||||
func (c *TCPCoalescer) commitParsed(pkt []byte, info *parsedTCP) error {
|
|
||||||
// Admission: only ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE may ride a coalesce chain. CWR marks a
|
|
||||||
// one-shot congestion transition the receiver must observe at a segment boundary. NB: AccECN
|
|
||||||
// reuses CWR as ACE counter bits; revisit this check if inner hosts adopt AccECN.
|
|
||||||
if info.flags&tcpFlagAck == 0 || info.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
|
|
||||||
// SYN/FIN/RST/URG/CWR must be observed in sequence. Seal the flow's open slot so later
|
|
||||||
// in-flow packets cannot extend it and emit ahead of this verbatim.
|
|
||||||
c.sealFlow(info.fk)
|
|
||||||
c.addVerbatim(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
if info.payLen == 0 {
|
|
||||||
// Pure ACK: no ordering obligation toward the flow's data. Delivering it after
|
|
||||||
// later-transmitted data only makes it a stale ACK, which receivers ignore. Not sealing
|
|
||||||
// keeps a bidirectional flow's data run coalescing across interleaved peer ACKs, matching
|
|
||||||
// kernel GRO. This is the only place emission deviates from transmission order.
|
|
||||||
c.addVerbatim(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Cached-slot fast path. Arrival isn't per-packet interleaved even with
|
|
||||||
// many flows: wire-side GRO delivers runs of same-flow packets
|
|
||||||
// (deliverSegments splits a superdatagram into up to 64), so the cache
|
|
||||||
// hits for the length of each run and a miss costs one fk compare
|
|
||||||
// before the map lookup carries the weight.
|
|
||||||
var open *coalesceSlot
|
|
||||||
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
|
||||||
open = last
|
|
||||||
} else {
|
|
||||||
open = c.openSlots[info.fk]
|
|
||||||
}
|
|
||||||
if open != nil {
|
|
||||||
if c.canAppend(open, pkt, info) {
|
|
||||||
if c.appendPayload(open, pkt, info) {
|
|
||||||
// Chain closed (PSH or short segment): stop extending it.
|
|
||||||
c.sealFlow(info.fk)
|
|
||||||
} else {
|
|
||||||
c.lastSlot = open
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
// Can't extend (seq gap from upstream loss, header change, or a full
|
|
||||||
// chain): evict it from openSlots and fall through to seed a fresh slot.
|
|
||||||
c.sealFlow(info.fk)
|
|
||||||
}
|
|
||||||
c.seed(pkt, info)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *TCPCoalescer) Flush() error {
|
|
||||||
var first error
|
|
||||||
for _, s := range c.slots {
|
|
||||||
var err error
|
|
||||||
if s.verbatim || s.numSeg == 1 {
|
|
||||||
// A slot that never grew is byte-identical to its seed packet; ship the original so
|
|
||||||
// its valid checksum rides the DATA_VALID path instead of a kernel software csum.
|
|
||||||
// rawPkt is only mutated once numSeg >= 2 (PSH propagate, flush patches), so it is
|
|
||||||
// pristine here.
|
|
||||||
_, err = c.w.Write(s.rawPkt)
|
|
||||||
} else {
|
|
||||||
err = c.flushSlot(s)
|
|
||||||
}
|
|
||||||
if err != nil && first == nil {
|
|
||||||
first = err
|
|
||||||
}
|
|
||||||
c.release(s)
|
|
||||||
}
|
|
||||||
clear(c.slots)
|
|
||||||
c.slots = c.slots[:0]
|
|
||||||
clear(c.openSlots)
|
|
||||||
c.lastSlot = nil
|
|
||||||
|
|
||||||
return first
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *TCPCoalescer) addVerbatim(pkt []byte) {
|
|
||||||
s := c.take()
|
|
||||||
s.verbatim = true
|
|
||||||
s.rawPkt = pkt
|
|
||||||
c.slots = append(c.slots, s)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *TCPCoalescer) seed(pkt []byte, info *parsedTCP) {
|
|
||||||
if info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
|
||||||
// Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
|
|
||||||
// be open here (commitParsed evicts before seeding), so sealFlow is defense in depth
|
|
||||||
// against a stale cache entry absorbing later data.
|
|
||||||
c.sealFlow(info.fk)
|
|
||||||
c.addVerbatim(pkt)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
s := c.take()
|
|
||||||
s.verbatim = false
|
|
||||||
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
|
|
||||||
// the superpacket header flushSlot patches in place.
|
|
||||||
s.rawPkt = pkt
|
|
||||||
s.hdrLen = info.hdrLen
|
|
||||||
s.ipHdrLen = info.ipHdrLen
|
|
||||||
s.isV6 = info.fk.isV6
|
|
||||||
s.fk = info.fk
|
|
||||||
s.gsoSize = info.payLen
|
|
||||||
s.numSeg = 1
|
|
||||||
s.totalPay = info.payLen
|
|
||||||
s.nextSeq = info.seq + uint32(info.payLen)
|
|
||||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
|
||||||
c.slots = append(c.slots, s)
|
|
||||||
if info.flags&tcpFlagPsh == 0 {
|
|
||||||
c.openSlots[info.fk] = s
|
|
||||||
c.lastSlot = s
|
|
||||||
} else {
|
|
||||||
// PSH on the seed closes the chain immediately; it is never registered as open.
|
|
||||||
// Drop any stale entry for this flow too (defense in depth, unreachable if lastSlot's lockstep invariant holds).
|
|
||||||
c.sealFlow(info.fk)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// canAppend reports whether info's packet extends the slot's seed: same header shape and stable
|
|
||||||
// contents, adjacent seq, not oversized. A closed chain never reaches here; closing removes the
|
|
||||||
// slot from openSlots, the only path in. The header fields read from rawPkt are always pristine:
|
|
||||||
// the only pre-flush mutation is the PSH propagate, which also closes the chain.
|
|
||||||
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
|
|
||||||
if info.hdrLen != s.hdrLen {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if info.seq != s.nextSeq {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if s.numSeg >= tcpCoalesceMaxSegs {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if info.payLen > s.gsoSize {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
// ECE state must be stable across a burst.
|
|
||||||
// Receivers expect the flag set on every segment of a CE-echoing window or none.
|
|
||||||
seedFlags := s.rawPkt[s.ipHdrLen+13]
|
|
||||||
if (seedFlags^info.flags)&tcpFlagEce != 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !headersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// appendPayload folds info's packet into s and reports whether the chain is now closed: the
|
|
||||||
// segment was sub-gsoSize (kernel TSO allows only the final segment to be short) or carried PSH.
|
|
||||||
// The caller must deregister a closed slot from openSlots.
|
|
||||||
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
|
|
||||||
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
|
||||||
s.numSeg++
|
|
||||||
s.totalPay += info.payLen
|
|
||||||
s.nextSeq = info.seq + uint32(info.payLen)
|
|
||||||
if info.flags&tcpFlagPsh != 0 {
|
|
||||||
// Propagate PSH into the seed header so kernel TSO sets it on the last segment. Mutating
|
|
||||||
// rawPkt is safe: PSH also closes the chain, so no admission check re-reads this header.
|
|
||||||
s.rawPkt[s.ipHdrLen+13] |= tcpFlagPsh
|
|
||||||
}
|
|
||||||
return info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *TCPCoalescer) take() *coalesceSlot {
|
|
||||||
if n := len(c.pool); n > 0 {
|
|
||||||
s := c.pool[n-1]
|
|
||||||
c.pool[n-1] = nil
|
|
||||||
c.pool = c.pool[:n-1]
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
return &coalesceSlot{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
|
||||||
clear(s.payIovs)
|
|
||||||
*s = coalesceSlot{payIovs: s.payIovs[:0]}
|
|
||||||
c.pool = append(c.pool, s)
|
|
||||||
}
|
|
||||||
|
|
||||||
// flushSlot patches the superpacket header in place in rawPkt (total length, IPv4 header
|
|
||||||
// checksum, pseudo-header checksum seed) and calls WriteGSO. The slot is released right after,
|
|
||||||
// so nothing re-reads the patched header. Does not remove the slot from c.slots.
|
|
||||||
func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
|
||||||
total := s.hdrLen + s.totalPay
|
|
||||||
l4Len := total - s.ipHdrLen
|
|
||||||
hdr := s.rawPkt[:s.hdrLen]
|
|
||||||
|
|
||||||
if s.isV6 {
|
|
||||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
|
||||||
} else {
|
|
||||||
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
|
|
||||||
hdr[10] = 0
|
|
||||||
hdr[11] = 0
|
|
||||||
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
|
|
||||||
}
|
|
||||||
|
|
||||||
var psum uint32
|
|
||||||
if s.isV6 {
|
|
||||||
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoTCP, l4Len)
|
|
||||||
} else {
|
|
||||||
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoTCP, l4Len)
|
|
||||||
}
|
|
||||||
tcsum := s.ipHdrLen + 16
|
|
||||||
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
|
||||||
|
|
||||||
return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
|
|
||||||
}
|
|
||||||
|
|
||||||
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
|
|
||||||
// equality on every field that must be identical across coalesced
|
|
||||||
// segments. Size/IPID/IPCsum/seq/flags/tcpCsum are masked out.
|
|
||||||
func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
|
||||||
if len(a) != len(b) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !ipHeadersMatch(a, b, isV6) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
// TCP: compare [0:4] ports, [8:13] ack+dataoff, [14:16] window,
|
|
||||||
// [18:tcpHdrLen] options (incl. urgent).
|
|
||||||
tcp := ipHdrLen
|
|
||||||
if !bytes.Equal(a[tcp:tcp+4], b[tcp:tcp+4]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !bytes.Equal(a[tcp+8:tcp+13], b[tcp+8:tcp+13]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !bytes.Equal(a[tcp+14:tcp+16], b[tcp+14:tcp+16]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !bytes.Equal(a[tcp+18:], b[tcp+18:]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
|
|
||||||
// already have its checksum field zeroed) and returns the folded/inverted
|
|
||||||
// 16-bit value to store.
|
|
||||||
func ipv4HdrChecksum(hdr []byte) uint16 {
|
|
||||||
var sum uint32
|
|
||||||
for i := 0; i+1 < len(hdr); i += 2 {
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(hdr[i : i+2]))
|
|
||||||
}
|
|
||||||
if len(hdr)%2 == 1 {
|
|
||||||
sum += uint32(hdr[len(hdr)-1]) << 8
|
|
||||||
}
|
|
||||||
for sum>>16 != 0 {
|
|
||||||
sum = (sum & 0xffff) + (sum >> 16)
|
|
||||||
}
|
|
||||||
return ^uint16(sum)
|
|
||||||
}
|
|
||||||
|
|
||||||
// pseudoSumIPv4 / pseudoSumIPv6 build the L4 pseudo-header partial sum
|
|
||||||
// expected by the virtio NEEDS_CSUM kernel path: the 32-bit accumulator
|
|
||||||
// before folding. proto selects the L4 (TCP or UDP); the UDP coalescer
|
|
||||||
// reuses these helpers.
|
|
||||||
func pseudoSumIPv4(src, dst []byte, proto byte, l4Len int) uint32 {
|
|
||||||
var sum uint32
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(src[0:2]))
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(src[2:4]))
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(dst[0:2]))
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(dst[2:4]))
|
|
||||||
sum += uint32(proto)
|
|
||||||
sum += uint32(l4Len)
|
|
||||||
return sum
|
|
||||||
}
|
|
||||||
|
|
||||||
func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
|
|
||||||
var sum uint32
|
|
||||||
for i := 0; i < 16; i += 2 {
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(src[i : i+2]))
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(dst[i : i+2]))
|
|
||||||
}
|
|
||||||
sum += uint32(l4Len >> 16)
|
|
||||||
sum += uint32(l4Len & 0xffff)
|
|
||||||
sum += uint32(proto)
|
|
||||||
return sum
|
|
||||||
}
|
|
||||||
|
|
||||||
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it unchanged (no one's complement).
|
|
||||||
// This is what virtio NEEDS_CSUM wants in the L4 checksum field
|
|
||||||
func foldOnceNoInvert(sum uint32) uint16 {
|
|
||||||
for sum>>16 != 0 {
|
|
||||||
sum = (sum & 0xffff) + (sum >> 16)
|
|
||||||
}
|
|
||||||
return uint16(sum)
|
|
||||||
}
|
|
||||||
@@ -1,214 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
)
|
|
||||||
|
|
||||||
// nopTunWriter is a zero-alloc tio.GSOWriter for benchmarks. Discards
|
|
||||||
// everything but satisfies the interface the coalescer detects.
|
|
||||||
type nopTunWriter struct{}
|
|
||||||
|
|
||||||
func (nopTunWriter) Write(p []byte) (int, error) { return len(p), nil }
|
|
||||||
func (nopTunWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, _ tio.GSOProto) error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
func (nopTunWriter) Capabilities() tio.Capabilities {
|
|
||||||
return tio.Capabilities{TSO: true, USO: true}
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildTCPv4BulkFlow returns a slice of N adjacent ACK-only TCP segments
|
|
||||||
// on a single 5-tuple, each carrying payloadLen bytes. Seq numbers are
|
|
||||||
// contiguous so every packet is coalesceable onto the previous one.
|
|
||||||
func buildTCPv4BulkFlow(n, payloadLen int) [][]byte {
|
|
||||||
pkts := make([][]byte, n)
|
|
||||||
pay := make([]byte, payloadLen)
|
|
||||||
seq := uint32(1000)
|
|
||||||
for i := range n {
|
|
||||||
pkts[i] = buildTCPv4(seq, tcpAck, pay)
|
|
||||||
seq += uint32(payloadLen)
|
|
||||||
}
|
|
||||||
return pkts
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildTCPv4Interleaved returns nFlows * perFlow packets with per-flow
|
|
||||||
// seq continuity but round-robin across flows — worst case for any
|
|
||||||
// "last-slot" cache.
|
|
||||||
func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte {
|
|
||||||
pay := make([]byte, payloadLen)
|
|
||||||
seqs := make([]uint32, nFlows)
|
|
||||||
for i := range seqs {
|
|
||||||
seqs[i] = uint32(1000 + i*1000000)
|
|
||||||
}
|
|
||||||
pkts := make([][]byte, 0, nFlows*perFlow)
|
|
||||||
for range perFlow {
|
|
||||||
for f := range nFlows {
|
|
||||||
sport := uint16(10000 + f)
|
|
||||||
pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay))
|
|
||||||
seqs[f] += uint32(payloadLen)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return pkts
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildTCPv4RunInterleaved returns nFlows*perFlow packets delivered in
|
|
||||||
// runs of runLen per flow — the arrival pattern wire-side GRO actually
|
|
||||||
// produces (deliverSegments splits each superdatagram into up to 64
|
|
||||||
// same-flow packets back to back). Contrast with buildTCPv4Interleaved's
|
|
||||||
// per-packet round-robin, the adversarial worst case for a last-slot cache.
|
|
||||||
func buildTCPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte {
|
|
||||||
pay := make([]byte, payloadLen)
|
|
||||||
seqs := make([]uint32, nFlows)
|
|
||||||
for i := range seqs {
|
|
||||||
seqs[i] = uint32(1000 + i*1000000)
|
|
||||||
}
|
|
||||||
pkts := make([][]byte, 0, nFlows*perFlow)
|
|
||||||
for done := 0; done < perFlow; done += runLen {
|
|
||||||
for f := range nFlows {
|
|
||||||
sport := uint16(10000 + f)
|
|
||||||
for range runLen {
|
|
||||||
pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay))
|
|
||||||
seqs[f] += uint32(payloadLen)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return pkts
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildICMPv4 returns a minimal non-TCP packet that takes the verbatim
|
|
||||||
// branch in Commit.
|
|
||||||
func buildICMPv4() []byte {
|
|
||||||
pkt := make([]byte, 28)
|
|
||||||
pkt[0] = 0x45
|
|
||||||
binary.BigEndian.PutUint16(pkt[2:4], 28)
|
|
||||||
pkt[9] = 1 // ICMP
|
|
||||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
|
||||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
// runCommitBench drives Commit over pkts batchSize at a time, flushing
|
|
||||||
// between batches, and reports per-packet cost.
|
|
||||||
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
|
||||||
b.Helper()
|
|
||||||
c := newTestTCPCoalescer(b, nopTunWriter{})
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.SetBytes(int64(len(pkts[0])))
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
pkt := pkts[i%len(pkts)]
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
if (i+1)%batchSize == 0 {
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// Drain any trailing partial batch so slot state doesn't leak across runs.
|
|
||||||
_ = c.Flush()
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkCommitSingleFlow is the bulk-TCP steady state: one flow,
|
|
||||||
// contiguous seq, 1200-byte payloads. Every packet past the seed should
|
|
||||||
// append onto the open slot. This is the case we most care about.
|
|
||||||
func BenchmarkCommitSingleFlow(b *testing.B) {
|
|
||||||
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
|
|
||||||
runCommitBench(b, pkts, tcpCoalesceMaxSegs)
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkCommitInterleaved4 has 4 concurrent bulk flows round-robined.
|
|
||||||
// A single-entry fast-path cache will miss on every packet; an N-way
|
|
||||||
// cache or map lookup carries the weight.
|
|
||||||
func BenchmarkCommitInterleaved4(b *testing.B) {
|
|
||||||
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
|
||||||
runCommitBench(b, pkts, len(pkts))
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkCommitInterleaved16 stresses the map at higher flow counts.
|
|
||||||
func BenchmarkCommitInterleaved16(b *testing.B) {
|
|
||||||
pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200)
|
|
||||||
runCommitBench(b, pkts, len(pkts))
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkCommitRunInterleaved4 is 4 concurrent flows arriving in
|
|
||||||
// GRO-burst runs of 16 — the realistic multi-flow pattern. A last-slot
|
|
||||||
// cache hits for the length of each run; the per-packet round-robin
|
|
||||||
// benches above are its worst case.
|
|
||||||
func BenchmarkCommitRunInterleaved4(b *testing.B) {
|
|
||||||
pkts := buildTCPv4RunInterleaved(4, tcpCoalesceMaxSegs, 16, 1200)
|
|
||||||
runCommitBench(b, pkts, len(pkts))
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseBase
|
|
||||||
// bails early and addVerbatim is the only work.
|
|
||||||
func BenchmarkCommitPassthrough(b *testing.B) {
|
|
||||||
pkt := buildICMPv4()
|
|
||||||
pkts := make([][]byte, 64)
|
|
||||||
for i := range pkts {
|
|
||||||
pkts[i] = pkt
|
|
||||||
}
|
|
||||||
runCommitBench(b, pkts, 64)
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow.
|
|
||||||
// Each packet takes the "TCP but not admissible" branch which does a
|
|
||||||
// map delete + verbatim. Measures the seal-without-slot cost.
|
|
||||||
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
|
||||||
pay := make([]byte, 0)
|
|
||||||
pkts := make([][]byte, 64)
|
|
||||||
for i := range pkts {
|
|
||||||
pkts[i] = buildTCPv4(uint32(1000+i), tcpSyn|tcpAck, pay)
|
|
||||||
}
|
|
||||||
runCommitBench(b, pkts, 64)
|
|
||||||
}
|
|
||||||
|
|
||||||
// runMultiCommitBench drives MultiCoalescer.Commit with in-order keys, so
|
|
||||||
// it includes the staging sort's already-sorted fast path plus the
|
|
||||||
// dispatch-time parse — the full steady-state cost of the batcher. The
|
|
||||||
// ParsedPackets are precomputed: in production they fall out of the
|
|
||||||
// firewall's newPacket, which this bench does not model.
|
|
||||||
func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
|
||||||
b.Helper()
|
|
||||||
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
|
|
||||||
pps := make([]*firewall.ParsedPacket, len(pkts))
|
|
||||||
for i, p := range pkts {
|
|
||||||
pps[i] = testPP(p)
|
|
||||||
}
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.SetBytes(int64(len(pkts[0])))
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
j := i % len(pkts)
|
|
||||||
if err := m.Commit(pkts[j], SortKey{Epoch: 1, Counter: uint64(i + 1)}, pps[j]); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
if (i+1)%batchSize == 0 {
|
|
||||||
if err := m.Flush(); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
_ = m.Flush()
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkMultiCommitSingleFlow is the multi-lane analogue of
|
|
||||||
// BenchmarkCommitSingleFlow — same workload but routed through the
|
|
||||||
// dispatcher. The delta vs the single-lane bench measures dispatcher
|
|
||||||
// overhead.
|
|
||||||
func BenchmarkMultiCommitSingleFlow(b *testing.B) {
|
|
||||||
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
|
|
||||||
runMultiCommitBench(b, pkts, tcpCoalesceMaxSegs)
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkMultiCommitInterleaved4 mirrors BenchmarkCommitInterleaved4
|
|
||||||
// through the dispatcher.
|
|
||||||
func BenchmarkMultiCommitInterleaved4(b *testing.B) {
|
|
||||||
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
|
||||||
runMultiCommitBench(b, pkts, len(pkts))
|
|
||||||
}
|
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -1,59 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import "net/netip"
|
|
||||||
|
|
||||||
const SendBatchCap = 128
|
|
||||||
|
|
||||||
// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush.
|
|
||||||
type batchWriter interface {
|
|
||||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
|
|
||||||
}
|
|
||||||
|
|
||||||
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
|
|
||||||
// One SendBatch is owned by each listenIn goroutine; no locking is needed.
|
|
||||||
// Slots are backed by an Arena (see its docs)
|
|
||||||
type SendBatch struct {
|
|
||||||
out batchWriter
|
|
||||||
bufs [][]byte
|
|
||||||
dsts []netip.AddrPort
|
|
||||||
arena *Arena
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewSendBatch makes a SendBatch with batchCap slots and an arenaSize byte buffer for slices to back those slots
|
|
||||||
func NewSendBatch(out batchWriter, batchCap, arenaSize int) *SendBatch {
|
|
||||||
return &SendBatch{
|
|
||||||
out: out,
|
|
||||||
bufs: make([][]byte, 0, batchCap),
|
|
||||||
dsts: make([]netip.AddrPort, 0, batchCap),
|
|
||||||
arena: NewArena(arenaSize),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (b *SendBatch) Reserve(sz int) []byte {
|
|
||||||
return b.arena.Reserve(sz)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Len reports how many packets are queued for the next Flush. Callers use
|
|
||||||
// it to flush incrementally once a full sendmmsg batch has accumulated,
|
|
||||||
// bounding how long the first packet of a large read batch waits.
|
|
||||||
func (b *SendBatch) Len() int { return len(b.bufs) }
|
|
||||||
|
|
||||||
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort) {
|
|
||||||
b.bufs = append(b.bufs, pkt)
|
|
||||||
b.dsts = append(b.dsts, dst)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Flush writes every queued packet and reports how many actually went out. A short count means some destinations
|
|
||||||
// were undeliverable; the batch is drained either way.
|
|
||||||
func (b *SendBatch) Flush() (int, error) {
|
|
||||||
var err error
|
|
||||||
written := 0
|
|
||||||
if len(b.bufs) > 0 {
|
|
||||||
written, err = b.out.WriteBatch(b.bufs, b.dsts)
|
|
||||||
}
|
|
||||||
clear(b.bufs)
|
|
||||||
b.bufs = b.bufs[:0]
|
|
||||||
b.dsts = b.dsts[:0]
|
|
||||||
b.arena.Reset()
|
|
||||||
return written, err
|
|
||||||
}
|
|
||||||
@@ -1,122 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
type fakeBatchWriter struct {
|
|
||||||
bufs [][]byte
|
|
||||||
addrs []netip.AddrPort
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
|
||||||
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
|
|
||||||
// returns, so tests must capture data before that happens.
|
|
||||||
w.bufs = make([][]byte, len(bufs))
|
|
||||||
for i, b := range bufs {
|
|
||||||
cp := make([]byte, len(b))
|
|
||||||
copy(cp, b)
|
|
||||||
w.bufs[i] = cp
|
|
||||||
}
|
|
||||||
w.addrs = append(w.addrs[:0], addrs...)
|
|
||||||
return len(bufs), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSendBatchReserveCommitFlush(t *testing.T) {
|
|
||||||
fw := &fakeBatchWriter{}
|
|
||||||
b := NewSendBatch(fw, 4, 32)
|
|
||||||
|
|
||||||
ap := netip.MustParseAddrPort("10.0.0.1:4242")
|
|
||||||
for i := 0; i < 4; i++ {
|
|
||||||
slot := b.Reserve(32)
|
|
||||||
if cap(slot) != 32 {
|
|
||||||
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
|
|
||||||
}
|
|
||||||
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
|
|
||||||
b.Commit(pkt, ap)
|
|
||||||
}
|
|
||||||
if _, err := b.Flush(); err != nil {
|
|
||||||
t.Fatalf("Flush: %v", err)
|
|
||||||
}
|
|
||||||
if len(fw.bufs) != 4 {
|
|
||||||
t.Fatalf("WriteBatch got %d bufs want 4", len(fw.bufs))
|
|
||||||
}
|
|
||||||
for i, buf := range fw.bufs {
|
|
||||||
if len(buf) != 3 || buf[0] != byte(i) {
|
|
||||||
t.Errorf("buf %d: %x", i, buf)
|
|
||||||
}
|
|
||||||
if fw.addrs[i] != ap {
|
|
||||||
t.Errorf("addr %d: got %v want %v", i, fw.addrs[i], ap)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Flush again with nothing committed — should be a no-op.
|
|
||||||
fw.bufs = nil
|
|
||||||
if _, err := b.Flush(); err != nil {
|
|
||||||
t.Fatalf("empty Flush: %v", err)
|
|
||||||
}
|
|
||||||
if fw.bufs != nil {
|
|
||||||
t.Fatalf("empty Flush triggered WriteBatch")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Reuse after Flush.
|
|
||||||
slot := b.Reserve(32)
|
|
||||||
if cap(slot) != 32 {
|
|
||||||
t.Fatalf("after Flush Reserve wrong cap: %d", cap(slot))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
|
||||||
fw := &fakeBatchWriter{}
|
|
||||||
b := NewSendBatch(fw, 3, 8)
|
|
||||||
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
|
||||||
|
|
||||||
for i := 0; i < 3; i++ {
|
|
||||||
s := b.Reserve(8)
|
|
||||||
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
|
|
||||||
b.Commit(pkt, ap)
|
|
||||||
}
|
|
||||||
if _, err := b.Flush(); err != nil {
|
|
||||||
t.Fatalf("Flush: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, buf := range fw.bufs {
|
|
||||||
if buf[0] != byte(0xA0+i) || buf[1] != byte(0xB0+i) {
|
|
||||||
t.Errorf("slot %d corrupted: %x", i, buf)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSendBatchGrowPreservesCommitted(t *testing.T) {
|
|
||||||
fw := &fakeBatchWriter{}
|
|
||||||
// Tiny initial backing forces a grow on the second Reserve.
|
|
||||||
b := NewSendBatch(fw, 1, 4)
|
|
||||||
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
|
||||||
|
|
||||||
s1 := b.Reserve(4)
|
|
||||||
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
|
|
||||||
b.Commit(pkt1, ap)
|
|
||||||
|
|
||||||
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
|
|
||||||
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
|
|
||||||
b.Commit(pkt2, ap)
|
|
||||||
|
|
||||||
// pkt1 must still be intact even though backing reallocated.
|
|
||||||
if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
|
|
||||||
t.Fatalf("first packet corrupted by grow: %x", pkt1)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := b.Flush(); err != nil {
|
|
||||||
t.Fatalf("Flush: %v", err)
|
|
||||||
}
|
|
||||||
if len(fw.bufs) != 2 {
|
|
||||||
t.Fatalf("got %d bufs want 2", len(fw.bufs))
|
|
||||||
}
|
|
||||||
if fw.bufs[0][0] != 0x11 || fw.bufs[0][3] != 0x44 {
|
|
||||||
t.Errorf("first packet on the wire: %x", fw.bufs[0])
|
|
||||||
}
|
|
||||||
if fw.bufs[1][0] != 0xA || fw.bufs[1][4] != 0xE {
|
|
||||||
t.Errorf("second packet on the wire: %x", fw.bufs[1])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,345 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"encoding/binary"
|
|
||||||
"io"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
)
|
|
||||||
|
|
||||||
// ipProtoUDP is the IANA protocol number for UDP.
|
|
||||||
const ipProtoUDP = 17
|
|
||||||
|
|
||||||
// udpCoalesceBufSize caps total bytes per UDP superpacket. Mirrors the
|
|
||||||
// kernel's gso_max_size; payloads beyond this are emitted as-is.
|
|
||||||
const udpCoalesceBufSize = 65535
|
|
||||||
|
|
||||||
// udpCoalesceMaxSegs caps how many segments we'll coalesce. Kernel UDP-GSO
|
|
||||||
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
|
|
||||||
const udpCoalesceMaxSegs = 64
|
|
||||||
|
|
||||||
// udpSlot is one entry in the UDPCoalescer's ordered event queue.
|
|
||||||
type udpSlot struct {
|
|
||||||
verbatim bool
|
|
||||||
// rawPkt is borrowed: the whole packet for verbatim slots, the seed
|
|
||||||
// packet for coalesce slots. A coalesce slot that never grows past one
|
|
||||||
// segment is emitted from rawPkt so its original (already valid) L4
|
|
||||||
// checksum ships DATA_VALID instead of making the kernel recompute it.
|
|
||||||
// A multi-segment slot's superpacket header is rawPkt's, patched in place at flush.
|
|
||||||
rawPkt []byte
|
|
||||||
|
|
||||||
fk flowKey
|
|
||||||
hdrLen int
|
|
||||||
ipHdrLen int
|
|
||||||
isV6 bool
|
|
||||||
gsoSize int // per-segment UDP payload length
|
|
||||||
numSeg int
|
|
||||||
totalPay int
|
|
||||||
payIovs [][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple
|
|
||||||
// concurrent flows and emits each flow's run as a single GSO_UDP_L4 superpacket via tio.GSOWriter.
|
|
||||||
// Preserves the in-flow order of packets as they are Commit-ed
|
|
||||||
//
|
|
||||||
// Owns no locks; one coalescer per TUN write queue.
|
|
||||||
type UDPCoalescer struct {
|
|
||||||
w tio.GSOWriter
|
|
||||||
slots []*udpSlot
|
|
||||||
openSlots map[flowKey]*udpSlot
|
|
||||||
// lastSlot caches the most recently touched open slot; see the
|
|
||||||
// TCPCoalescer field of the same name. Single-flow QUIC bulk is the
|
|
||||||
// dominant USO workload, and multi-flow arrival comes in GRO runs, so
|
|
||||||
// the fk compare beats the map's 38-byte key hash on most packets.
|
|
||||||
// Kept in lockstep with openSlots: nil whenever the slot it pointed at
|
|
||||||
// is removed.
|
|
||||||
lastSlot *udpSlot
|
|
||||||
pool []*udpSlot
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewUDPCoalescer(w io.Writer) *UDPCoalescer {
|
|
||||||
gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP)
|
|
||||||
if !ok {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return &UDPCoalescer{
|
|
||||||
w: gw,
|
|
||||||
slots: make([]*udpSlot, 0, initialSlots),
|
|
||||||
openSlots: make(map[flowKey]*udpSlot, initialSlots),
|
|
||||||
pool: make([]*udpSlot, 0, initialSlots),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// parsedUDP holds the fields extracted from a single parse so later steps
|
|
||||||
// (admission, slot lookup, canAppend) don't re-walk the header.
|
|
||||||
type parsedUDP struct {
|
|
||||||
fk flowKey
|
|
||||||
ipHdrLen int
|
|
||||||
hdrLen int // ipHdrLen + 8
|
|
||||||
payLen int
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseAt extracts the flow key and IP/UDP offsets for a packet the dispatcher already knows is
|
|
||||||
// UDP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
|
|
||||||
// entry and is filled in place. Returns false for malformed input or any shape that must not
|
|
||||||
// coalesce (IPv4 options/fragmentation, IPv6 extension headers).
|
|
||||||
func (p *parsedUDP) parseAt(pkt []byte, ipHdrLen int) bool {
|
|
||||||
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
|
|
||||||
if !ok {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return p.parseTail(trimmed, ipHdrLen)
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseTail layers the UDP-header parse on a validated IP prologue. pkt is the trimmed packet;
|
|
||||||
// fk's addresses are already filled.
|
|
||||||
func (p *parsedUDP) parseTail(pkt []byte, ipHdrLen int) bool {
|
|
||||||
if len(pkt) < ipHdrLen+8 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
|
|
||||||
udpLen := int(binary.BigEndian.Uint16(pkt[ipHdrLen+4 : ipHdrLen+6]))
|
|
||||||
if udpLen < 8 || udpLen > len(pkt)-ipHdrLen {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
p.ipHdrLen = ipHdrLen
|
|
||||||
p.hdrLen = ipHdrLen + 8
|
|
||||||
p.payLen = udpLen - 8
|
|
||||||
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
|
|
||||||
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
|
|
||||||
// hashing the 38-byte key when no chains are open.
|
|
||||||
func (c *UDPCoalescer) sealFlow(fk flowKey) {
|
|
||||||
if len(c.openSlots) == 0 {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if last := c.lastSlot; last != nil && last.fk == fk {
|
|
||||||
c.lastSlot = nil
|
|
||||||
}
|
|
||||||
delete(c.openSlots, fk)
|
|
||||||
}
|
|
||||||
|
|
||||||
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
|
|
||||||
// coalesce (any fragmentation, unparseable header) seals every open chain — its flow is unknown —
|
|
||||||
// and rides the lane as an in-lane verbatim, still in transmission order.
|
|
||||||
func (c *UDPCoalescer) commitStaged(sp stagedPacket) error {
|
|
||||||
if sp.fragAny {
|
|
||||||
c.sealAllOpen()
|
|
||||||
c.addVerbatim(sp.pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
var info parsedUDP
|
|
||||||
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
|
||||||
c.sealAllOpen()
|
|
||||||
c.addVerbatim(sp.pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return c.commitParsed(sp.pkt, &info)
|
|
||||||
}
|
|
||||||
|
|
||||||
// commitParsed commits one parsed UDP packet. The caller (dispatch, via parseAt) supplies a
|
|
||||||
// valid parse so the header is not re-walked here.
|
|
||||||
func (c *UDPCoalescer) commitParsed(pkt []byte, info *parsedUDP) error {
|
|
||||||
// A zero-length UDP datagram (length == 8) is legal and must reach the TUN, but cannot be
|
|
||||||
// coalesced.
|
|
||||||
if info.payLen == 0 {
|
|
||||||
c.sealFlow(info.fk)
|
|
||||||
c.addVerbatim(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
// Cached-slot fast path; see the TCPCoalescer equivalent.
|
|
||||||
var open *udpSlot
|
|
||||||
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
|
||||||
open = last
|
|
||||||
} else {
|
|
||||||
open = c.openSlots[info.fk]
|
|
||||||
}
|
|
||||||
if open != nil {
|
|
||||||
if c.canAppend(open, pkt, info) {
|
|
||||||
if c.appendPayload(open, pkt, info) {
|
|
||||||
// Chain closed (short segment): stop extending it.
|
|
||||||
c.sealFlow(info.fk)
|
|
||||||
} else {
|
|
||||||
c.lastSlot = open
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
// Can't extend: evict it from openSlots and fall through to seed a
|
|
||||||
// fresh slot.
|
|
||||||
c.sealFlow(info.fk)
|
|
||||||
}
|
|
||||||
c.seed(pkt, info)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) Flush() error {
|
|
||||||
var first error
|
|
||||||
for _, s := range c.slots {
|
|
||||||
var err error
|
|
||||||
if s.verbatim || s.numSeg == 1 {
|
|
||||||
// A slot that never grew is byte-identical to the packet it was
|
|
||||||
// seeded from; ship the original so its valid checksum rides the
|
|
||||||
// DATA_VALID path instead of paying a kernel software csum.
|
|
||||||
_, err = c.w.Write(s.rawPkt)
|
|
||||||
} else {
|
|
||||||
err = c.flushSlot(s)
|
|
||||||
}
|
|
||||||
if err != nil && first == nil {
|
|
||||||
first = err
|
|
||||||
}
|
|
||||||
c.release(s)
|
|
||||||
}
|
|
||||||
clear(c.slots)
|
|
||||||
c.slots = c.slots[:0]
|
|
||||||
clear(c.openSlots)
|
|
||||||
c.lastSlot = nil
|
|
||||||
return first
|
|
||||||
}
|
|
||||||
|
|
||||||
// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
|
|
||||||
// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
|
|
||||||
func (c *UDPCoalescer) sealAllOpen() {
|
|
||||||
clear(c.openSlots)
|
|
||||||
c.lastSlot = nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) addVerbatim(pkt []byte) {
|
|
||||||
s := c.take()
|
|
||||||
s.verbatim = true
|
|
||||||
s.rawPkt = pkt
|
|
||||||
c.slots = append(c.slots, s)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) seed(pkt []byte, info *parsedUDP) {
|
|
||||||
if info.hdrLen+info.payLen > udpCoalesceBufSize {
|
|
||||||
// Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
|
|
||||||
// be open here (commitParsed evicts before seeding), so sealFlow is defense in depth
|
|
||||||
// against a stale cache entry absorbing later data.
|
|
||||||
c.sealFlow(info.fk)
|
|
||||||
c.addVerbatim(pkt)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
s := c.take()
|
|
||||||
s.verbatim = false
|
|
||||||
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
|
|
||||||
// the superpacket header flushSlot patches in place.
|
|
||||||
s.rawPkt = pkt
|
|
||||||
s.hdrLen = info.hdrLen
|
|
||||||
s.ipHdrLen = info.ipHdrLen
|
|
||||||
s.isV6 = info.fk.isV6
|
|
||||||
s.fk = info.fk
|
|
||||||
s.gsoSize = info.payLen
|
|
||||||
s.numSeg = 1
|
|
||||||
s.totalPay = info.payLen
|
|
||||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
|
||||||
c.slots = append(c.slots, s)
|
|
||||||
c.openSlots[info.fk] = s
|
|
||||||
c.lastSlot = s
|
|
||||||
}
|
|
||||||
|
|
||||||
// canAppend reports whether info's packet extends the slot's seed.
|
|
||||||
// Kernel UDP-GSO requires every segment except possibly the last to be
|
|
||||||
// exactly gsoSize, and the last may be shorter (≤ gsoSize).
|
|
||||||
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info *parsedUDP) bool {
|
|
||||||
if info.hdrLen != s.hdrLen {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if s.numSeg >= udpCoalesceMaxSegs {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if info.payLen > s.gsoSize {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
// Header reads use rawPkt, which is never mutated before flush. A closed chain never reaches
|
|
||||||
// here; closing removes the slot from openSlots, the only path in.
|
|
||||||
if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !udpHeadersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// appendPayload folds info's packet into s and reports whether the chain is now closed: kernel
|
|
||||||
// UDP-GSO requires every segment but the last to be exactly gsoSize, so a short segment must be
|
|
||||||
// the final one. The caller must deregister a closed slot from openSlots.
|
|
||||||
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info *parsedUDP) bool {
|
|
||||||
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
|
||||||
s.numSeg++
|
|
||||||
s.totalPay += info.payLen
|
|
||||||
return info.payLen < s.gsoSize
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) take() *udpSlot {
|
|
||||||
if n := len(c.pool); n > 0 {
|
|
||||||
s := c.pool[n-1]
|
|
||||||
c.pool[n-1] = nil
|
|
||||||
c.pool = c.pool[:n-1]
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
return &udpSlot{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) release(s *udpSlot) {
|
|
||||||
// Reset every field, identity ones included; see TCPCoalescer.release.
|
|
||||||
clear(s.payIovs)
|
|
||||||
*s = udpSlot{payIovs: s.payIovs[:0]}
|
|
||||||
c.pool = append(c.pool, s)
|
|
||||||
}
|
|
||||||
|
|
||||||
// flushSlot patches the IP header total length / IPv6 payload length and
|
|
||||||
// the UDP length to the *total* across all coalesced segments, then seeds
|
|
||||||
// the UDP checksum field with the pseudo-header partial (single-fold, not
|
|
||||||
// inverted) per virtio NEEDS_CSUM. The patches land in place in rawPkt; the
|
|
||||||
// slot is released right after, so nothing re-reads the patched header.
|
|
||||||
func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
|
||||||
hdr := s.rawPkt[:s.hdrLen]
|
|
||||||
total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes
|
|
||||||
l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads)
|
|
||||||
|
|
||||||
if s.isV6 {
|
|
||||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
|
||||||
} else {
|
|
||||||
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
|
|
||||||
hdr[10] = 0
|
|
||||||
hdr[11] = 0
|
|
||||||
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
|
|
||||||
}
|
|
||||||
|
|
||||||
// UDP length field (offset 4 inside the UDP header) = total UDP size.
|
|
||||||
binary.BigEndian.PutUint16(hdr[s.ipHdrLen+4:s.ipHdrLen+6], uint16(l4Len))
|
|
||||||
|
|
||||||
var psum uint32
|
|
||||||
if s.isV6 {
|
|
||||||
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoUDP, l4Len)
|
|
||||||
} else {
|
|
||||||
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoUDP, l4Len)
|
|
||||||
}
|
|
||||||
udpCsumOff := s.ipHdrLen + 6
|
|
||||||
binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum))
|
|
||||||
|
|
||||||
return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP)
|
|
||||||
}
|
|
||||||
|
|
||||||
// udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on
|
|
||||||
// every field that must be identical across coalesced segments
|
|
||||||
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
|
||||||
if len(a) != len(b) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !ipHeadersMatch(a, b, isV6) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8]:
|
|
||||||
// length varies (we rewrite at flush) and the checksum will be redone.
|
|
||||||
udp := ipHdrLen
|
|
||||||
return bytes.Equal(a[udp:udp+4], b[udp:udp+4])
|
|
||||||
}
|
|
||||||
@@ -1,72 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
// buildUDPv4BulkFlow returns n equal-size datagrams on one flow — the
|
|
||||||
// steady state for single-flow QUIC bulk, the workload USO exists for.
|
|
||||||
func buildUDPv4BulkFlow(n, payloadLen int) [][]byte {
|
|
||||||
pay := make([]byte, payloadLen)
|
|
||||||
pkts := make([][]byte, n)
|
|
||||||
for i := range pkts {
|
|
||||||
pkts[i] = buildUDPv4(40000, 443, pay)
|
|
||||||
}
|
|
||||||
return pkts
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildUDPv4RunInterleaved mirrors buildTCPv4RunInterleaved: nFlows*perFlow
|
|
||||||
// datagrams arriving in GRO-burst runs of runLen per flow.
|
|
||||||
func buildUDPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte {
|
|
||||||
pay := make([]byte, payloadLen)
|
|
||||||
pkts := make([][]byte, 0, nFlows*perFlow)
|
|
||||||
for done := 0; done < perFlow; done += runLen {
|
|
||||||
for f := range nFlows {
|
|
||||||
sport := uint16(40000 + f)
|
|
||||||
for range runLen {
|
|
||||||
pkts = append(pkts, buildUDPv4(sport, 443, pay))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return pkts
|
|
||||||
}
|
|
||||||
|
|
||||||
// runUDPCommitBench drives UDPCoalescer.Commit over pkts batchSize at a
|
|
||||||
// time, flushing between batches, and reports per-packet cost.
|
|
||||||
func runUDPCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
|
||||||
b.Helper()
|
|
||||||
c := newTestUDPCoalescer(b, nopTunWriter{})
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.SetBytes(int64(len(pkts[0])))
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
pkt := pkts[i%len(pkts)]
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
if (i+1)%batchSize == 0 {
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
_ = c.Flush()
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkUDPCommitSingleFlow is the single-flow bulk steady state.
|
|
||||||
func BenchmarkUDPCommitSingleFlow(b *testing.B) {
|
|
||||||
pkts := buildUDPv4BulkFlow(udpCoalesceMaxSegs, 1200)
|
|
||||||
runUDPCommitBench(b, pkts, udpCoalesceMaxSegs)
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkUDPCommitInterleaved4 is the adversarial per-packet round-robin.
|
|
||||||
func BenchmarkUDPCommitInterleaved4(b *testing.B) {
|
|
||||||
pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 1, 1200)
|
|
||||||
runUDPCommitBench(b, pkts, len(pkts))
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkUDPCommitRunInterleaved4 is 4 flows in GRO-burst runs of 16.
|
|
||||||
func BenchmarkUDPCommitRunInterleaved4(b *testing.B) {
|
|
||||||
pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 16, 1200)
|
|
||||||
runUDPCommitBench(b, pkts, len(pkts))
|
|
||||||
}
|
|
||||||
@@ -1,536 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"encoding/binary"
|
|
||||||
"io"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
// buildUDPv4 builds a minimal IPv4+UDP packet with the given payload and ports.
|
|
||||||
func buildUDPv4(sport, dport uint16, payload []byte) []byte {
|
|
||||||
const ipHdrLen = 20
|
|
||||||
const udpHdrLen = 8
|
|
||||||
total := ipHdrLen + udpHdrLen + len(payload)
|
|
||||||
pkt := make([]byte, total)
|
|
||||||
|
|
||||||
pkt[0] = 0x45
|
|
||||||
pkt[1] = 0x00
|
|
||||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(total))
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], 0)
|
|
||||||
binary.BigEndian.PutUint16(pkt[6:8], 0x4000)
|
|
||||||
pkt[8] = 64
|
|
||||||
pkt[9] = ipProtoUDP
|
|
||||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
|
||||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(pkt[20:22], sport)
|
|
||||||
binary.BigEndian.PutUint16(pkt[22:24], dport)
|
|
||||||
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpHdrLen+len(payload)))
|
|
||||||
binary.BigEndian.PutUint16(pkt[26:28], 0)
|
|
||||||
|
|
||||||
copy(pkt[28:], payload)
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildUDPv6 builds a minimal IPv6+UDP packet.
|
|
||||||
func buildUDPv6(sport, dport uint16, payload []byte) []byte {
|
|
||||||
const ipHdrLen = 40
|
|
||||||
const udpHdrLen = 8
|
|
||||||
total := ipHdrLen + udpHdrLen + len(payload)
|
|
||||||
pkt := make([]byte, total)
|
|
||||||
|
|
||||||
pkt[0] = 0x60
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(udpHdrLen+len(payload)))
|
|
||||||
pkt[6] = ipProtoUDP
|
|
||||||
pkt[7] = 64
|
|
||||||
pkt[8] = 0xfe
|
|
||||||
pkt[9] = 0x80
|
|
||||||
pkt[23] = 1
|
|
||||||
pkt[24] = 0xfe
|
|
||||||
pkt[25] = 0x80
|
|
||||||
pkt[39] = 2
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(pkt[40:42], sport)
|
|
||||||
binary.BigEndian.PutUint16(pkt[42:44], dport)
|
|
||||||
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpHdrLen+len(payload)))
|
|
||||||
binary.BigEndian.PutUint16(pkt[46:48], 0)
|
|
||||||
|
|
||||||
copy(pkt[48:], payload)
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
// newTestUDPCoalescer builds a coalescer over w and fails the test if w can't
|
|
||||||
// do USO. See newTestTCPCoalescer.
|
|
||||||
func newTestUDPCoalescer(tb testing.TB, w io.Writer) *UDPCoalescer {
|
|
||||||
tb.Helper()
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
if c == nil {
|
|
||||||
tb.Fatal("NewUDPCoalescer: writer does not support USO")
|
|
||||||
}
|
|
||||||
return c
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNewUDPCoalescerRefusesWhenGSOUnavailable mirrors the TCP precondition:
|
|
||||||
// no USO, no coalescer.
|
|
||||||
func TestNewUDPCoalescerRefusesWhenGSOUnavailable(t *testing.T) {
|
|
||||||
if c := NewUDPCoalescer(&fakeTunWriter{gsoEnabled: false}); c != nil {
|
|
||||||
t.Fatalf("want nil for a non-USO writer, got %v", c)
|
|
||||||
}
|
|
||||||
if c := NewUDPCoalescer(&plainOnlyWriter{}); c != nil {
|
|
||||||
t.Fatalf("want nil for a plain writer, got %v", c)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
// ICMP packet
|
|
||||||
pkt := make([]byte, 28)
|
|
||||||
pkt[0] = 0x45
|
|
||||||
binary.BigEndian.PutUint16(pkt[2:4], 28)
|
|
||||||
pkt[9] = 1
|
|
||||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
|
||||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("ICMP must pass through unchanged: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pkt := buildUDPv4(1000, 53, make([]byte, 800))
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// A slot that never grew past one datagram flushes as a plain Write of
|
|
||||||
// the original packet bytes: the original (already valid) checksum
|
|
||||||
// ships via the DATA_VALID path, so the kernel does no csum work.
|
|
||||||
// WriteGSO is reserved for slots that actually coalesced (>=2 segs).
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if !bytes.Equal(w.writes[0], pkt) {
|
|
||||||
t.Errorf("plain write not byte-identical to committed packet: got %d bytes want %d", len(w.writes[0]), len(pkt))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
for i := 0; i < 3; i++ {
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
if g.gsoSize != 1200 {
|
|
||||||
t.Errorf("gsoSize=%d want 1200", g.gsoSize)
|
|
||||||
}
|
|
||||||
if len(g.pays) != 3 {
|
|
||||||
t.Errorf("pay count=%d want 3", len(g.pays))
|
|
||||||
}
|
|
||||||
if g.csumStart != 20 {
|
|
||||||
t.Errorf("csumStart=%d want 20", g.csumStart)
|
|
||||||
}
|
|
||||||
// IP totalLen and UDP length must be the TOTAL across all segments —
|
|
||||||
// the kernel's ip_rcv_core trims skbs to iph->tot_len, so a per-segment
|
|
||||||
// value would silently drop everything but the first segment. Total =
|
|
||||||
// IP(20) + UDP(8) + 3*1200 = 3628.
|
|
||||||
gotTotalLen := binary.BigEndian.Uint16(g.hdr[2:4])
|
|
||||||
if gotTotalLen != 3628 {
|
|
||||||
t.Errorf("ipv4 total_len=%d want 3628 (must be total across segments)", gotTotalLen)
|
|
||||||
}
|
|
||||||
gotUDPLen := binary.BigEndian.Uint16(g.hdr[20+4 : 20+6])
|
|
||||||
if gotUDPLen != 8+3*1200 {
|
|
||||||
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Last segment may be shorter, sealing the chain.
|
|
||||||
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
full := make([]byte, 1200)
|
|
||||||
tail := make([]byte, 600)
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, tail)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// A 4th packet, even same-sized, must NOT join — chain is sealed.
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// The sealed 3-datagram chain is a real superpacket; the re-seed stays
|
|
||||||
// single-segment and flushes as a plain write of the original packet.
|
|
||||||
if len(w.gsoWrites) != 1 || len(w.writes) != 1 {
|
|
||||||
t.Fatalf("want 1 gso (sealed) + 1 plain (new seed), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites[0].pays) != 3 {
|
|
||||||
t.Errorf("super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
|
||||||
}
|
|
||||||
if got, want := len(w.writes[0]), 20+8+1200; got != want {
|
|
||||||
t.Errorf("re-seed plain write len=%d want %d", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// A larger-than-gsoSize packet cannot extend the slot — it reseeds.
|
|
||||||
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 1200))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Both seeds stay single-segment → two plain writes in arrival order.
|
|
||||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("want 2 separate plain writes, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
for i, want := range []int{20 + 8 + 800, 20 + 8 + 1200} {
|
|
||||||
if len(w.writes[i]) != want {
|
|
||||||
t.Errorf("write %d len=%d want %d", i, len(w.writes[i]), want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Different 5-tuples must not coalesce.
|
|
||||||
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pay := make([]byte, 800)
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Two flows × 2 datagrams each = 2 superpackets of 2 segments.
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes (one per flow), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
for i, g := range w.gsoWrites {
|
|
||||||
if len(g.pays) != 2 {
|
|
||||||
t.Errorf("super %d: want 2 pays, got %d", i, len(g.pays))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Caps at udpCoalesceMaxSegs.
|
|
||||||
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pay := make([]byte, 100)
|
|
||||||
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// First superpacket holds udpCoalesceMaxSegs segments; the spillover
|
|
||||||
// reseeds a new one.
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes (cap then reseed), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites[0].pays) != udpCoalesceMaxSegs {
|
|
||||||
t.Errorf("first super: pays=%d want %d", len(w.gsoWrites[0].pays), udpCoalesceMaxSegs)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites[1].pays) != 5 {
|
|
||||||
t.Errorf("second super: pays=%d want 5", len(w.gsoWrites[1].pays))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Differing IP ECN codepoints must not coalesce: udpHeadersMatch compares
|
|
||||||
// the full ToS byte (matching kernel GRO). A CE-marked datagram mid-run
|
|
||||||
// seals the Not-ECT chain and reseeds; the trailing Not-ECT datagram
|
|
||||||
// reseeds again. All three stay single-segment, so each ships as a plain
|
|
||||||
// write of its original bytes, keeping its own codepoint.
|
|
||||||
func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pay := make([]byte, 800)
|
|
||||||
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT)
|
|
||||||
pkt1 := buildUDPv4(1000, 53, pay)
|
|
||||||
pkt1[1] = 0x03 // CE
|
|
||||||
pkt2 := buildUDPv4(1000, 53, pay) // ECN=00 again
|
|
||||||
for _, p := range [][]byte{pkt0, pkt1, pkt2} {
|
|
||||||
if err := c.Commit(p); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("want 3 separate plain writes (differing ECN), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
wantECN := []byte{0x00, 0x03, 0x00}
|
|
||||||
for i, p := range w.writes {
|
|
||||||
if got := p[1] & 0x03; got != wantECN[i] {
|
|
||||||
t.Errorf("write %d ECN=%#x want %#x", i, got, wantECN[i])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// IPv6 path: same flow, equal-sized → coalesced.
|
|
||||||
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
for i := 0; i < 3; i++ {
|
|
||||||
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 gso write, got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
if !g.isV6 {
|
|
||||||
t.Errorf("expected v6 write")
|
|
||||||
}
|
|
||||||
if g.csumStart != 40 {
|
|
||||||
t.Errorf("csumStart=%d want 40", g.csumStart)
|
|
||||||
}
|
|
||||||
// IPv6 payload_len and UDP length must be TOTAL — kernel's
|
|
||||||
// ip6_rcv_core trims to payload_len + ipv6 hdr size. Total UDP = 8 +
|
|
||||||
// 3*1200 = 3608.
|
|
||||||
gotPlen := binary.BigEndian.Uint16(g.hdr[4:6])
|
|
||||||
if gotPlen != 8+3*1200 {
|
|
||||||
t.Errorf("ipv6 payload_len=%d want %d (must be total)", gotPlen, 8+3*1200)
|
|
||||||
}
|
|
||||||
gotUDPLen := binary.BigEndian.Uint16(g.hdr[40+4 : 40+6])
|
|
||||||
if gotUDPLen != 8+3*1200 {
|
|
||||||
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte.
|
|
||||||
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pay := make([]byte, 800)
|
|
||||||
pkt0 := buildUDPv4(1000, 53, pay)
|
|
||||||
pkt1 := buildUDPv4(1000, 53, pay)
|
|
||||||
pkt1[1] = 0xb8 // EF DSCP, ECN=0
|
|
||||||
if err := c.Commit(pkt0); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(pkt1); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Both seeds stay single-segment → two plain writes, no gso.
|
|
||||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("want 2 separate plain writes (different DSCP), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Fragmented IPv4 must not be coalesced.
|
|
||||||
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
|
||||||
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("frag must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// A zero-length UDP datagram (UDP length == 8, no payload) is legal and
|
|
||||||
// must be delivered as a plain single datagram — never coalesced. Seeding
|
|
||||||
// it into a GSO slot stores an empty payload iovec that panics WriteGSO
|
|
||||||
// (index-out-of-range on &pay[0]); this is a remote DoS if we ever let it
|
|
||||||
// reach the GSO path. Regression: must not panic and must be written.
|
|
||||||
func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("zero-length UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.writes[0]) != len(pkt) {
|
|
||||||
t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// IPv6 zero-length UDP datagram: same verbatim contract as v4.
|
|
||||||
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("zero-length IPv6 UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.writes[0]) != len(pkt) {
|
|
||||||
t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// A zero-length datagram arriving mid-flow must seal the open chain so the
|
|
||||||
// datagram after it seeds a fresh superpacket *after* the empty one on the
|
|
||||||
// wire — per-flow arrival order (full, empty, full) must be preserved.
|
|
||||||
func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
full := make([]byte, 800)
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, nil)); err != nil { // zero-length
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// The empty datagram sealed the first slot, so the trailing full packet
|
|
||||||
// can't join it. All three emit as plain writes (the two full datagrams
|
|
||||||
// stayed single-segment; the empty one is verbatim) in per-flow
|
|
||||||
// arrival order: full, empty, full.
|
|
||||||
if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("want 3 plain writes, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
|
||||||
}
|
|
||||||
for i, want := range []int{20 + 8 + 800, 20 + 8, 20 + 8 + 800} {
|
|
||||||
if len(w.writes[i]) != want {
|
|
||||||
t.Errorf("write %d len=%d want %d (order full, empty, full)", i, len(w.writes[i]), want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// IPv4 with options is not admissible (we require IHL=5).
|
|
||||||
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
|
||||||
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("ipv4-with-options must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestUDPCoalescerNonAtomicSequentialIDsCoalesce mirrors the TCP rule: DF
|
|
||||||
// clear is fine as long as the IDs already run seed+1 per datagram, so
|
|
||||||
// kernel USO's re-stamp reproduces them.
|
|
||||||
func TestUDPCoalescerNonAtomicSequentialIDsCoalesce(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
|
|
||||||
for i := range 2 {
|
|
||||||
pkt := buildUDPv4(40000, 443, pay)
|
|
||||||
setIPv4ID(pkt, uint16(40+i), false)
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 || len(w.gsoWrites[0].pays) != 2 {
|
|
||||||
t.Fatalf("sequential-ID DF=0 datagrams must coalesce: gso=%d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestUDPCoalescerNonAtomicIDGapReseeds: an ID jump on a DF=0 flow breaks
|
|
||||||
// the chain; each datagram stays a single-segment slot and flushes as a
|
|
||||||
// plain write that keeps its own (meaningful) ID.
|
|
||||||
func TestUDPCoalescerNonAtomicIDGapReseeds(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := newTestUDPCoalescer(t, w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
|
|
||||||
p1 := buildUDPv4(40000, 443, pay)
|
|
||||||
setIPv4ID(p1, 40, false)
|
|
||||||
p2 := buildUDPv4(40000, 443, pay)
|
|
||||||
setIPv4ID(p2, 50, false)
|
|
||||||
|
|
||||||
if err := c.Commit(p1); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(p2); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("ID gap on DF=0 must reseed: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
for i, want := range []uint16{40, 50} {
|
|
||||||
if id := binary.BigEndian.Uint16(w.writes[i][4:6]); id != want {
|
|
||||||
t.Errorf("write %d: ID=%d want %d (must be preserved)", i, id, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,23 +0,0 @@
|
|||||||
package checksum
|
|
||||||
|
|
||||||
import (
|
|
||||||
"golang.org/x/sys/cpu"
|
|
||||||
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
|
||||||
)
|
|
||||||
|
|
||||||
//go:noescape
|
|
||||||
func checksumAVX2(buf []byte, initial uint16) uint16
|
|
||||||
|
|
||||||
var hasAVX2 = cpu.X86.HasAVX2
|
|
||||||
|
|
||||||
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
|
||||||
// initial. It is a drop-in replacement for gvisor's checksum.Checksum that
|
|
||||||
// dispatches to a hand-written AVX2 routine on amd64 CPUs that support it,
|
|
||||||
// falling back to gvisor's pure-Go implementation otherwise. The result
|
|
||||||
// matches gvisor's bit-for-bit for any buffer length and initial seed.
|
|
||||||
func Checksum(buf []byte, initial uint16) uint16 {
|
|
||||||
if hasAVX2 {
|
|
||||||
return checksumAVX2(buf, initial)
|
|
||||||
}
|
|
||||||
return gvisorchecksum.Checksum(buf, initial)
|
|
||||||
}
|
|
||||||
@@ -1,157 +0,0 @@
|
|||||||
#include "textflag.h"
|
|
||||||
|
|
||||||
// func checksumAVX2(buf []byte, initial uint16) uint16
|
|
||||||
//
|
|
||||||
// Computes the RFC 1071 ones-complement sum of buf, seeded with initial.
|
|
||||||
//
|
|
||||||
// Algorithm: sum the buffer treating it as a stream of uint32s in machine
|
|
||||||
// (little-endian) byte order, accumulating into 64-bit lanes (top 32 bits
|
|
||||||
// hold cross-add carries — at 1 byte / lane / iter we have 32 bits of
|
|
||||||
// headroom which is far more than the 16 KB/64 KB max practical inputs).
|
|
||||||
// At the end we fold to 16 bits and byte-swap once to recover the on-wire
|
|
||||||
// (big-endian) result. RFC 1071 §1.2.B byte-order independence makes this
|
|
||||||
// equivalent to summing as 16-bit big-endian words.
|
|
||||||
//
|
|
||||||
// The ymm accumulators (Y4..Y7) hold 4 uint64 lanes each = 16 parallel
|
|
||||||
// partial sums. The main loop loads 64 bytes per iter as four 16-byte
|
|
||||||
// chunks, zero-extending each chunk's four uint32s into a ymm via
|
|
||||||
// VPMOVZXDQ-from-memory, then VPADDQ into a separate accumulator per
|
|
||||||
// chunk to break the dep chain. After the vector loop the lane sums are
|
|
||||||
// horizontally reduced and merged with a scalar accumulator that handles
|
|
||||||
// the trailing 0..63 bytes plus the (byte-swapped) initial seed.
|
|
||||||
TEXT ·checksumAVX2(SB), NOSPLIT, $0-34
|
|
||||||
MOVQ buf_base+0(FP), SI
|
|
||||||
MOVQ buf_len+8(FP), CX
|
|
||||||
MOVWQZX initial+24(FP), AX
|
|
||||||
|
|
||||||
// Pre-byteswap initial into the LE-summing space so it merges directly
|
|
||||||
// with the rest of the accumulator. The final fold's bswap16 will undo
|
|
||||||
// this and convert the whole result back to BE.
|
|
||||||
XCHGB AH, AL
|
|
||||||
|
|
||||||
CMPQ CX, $32
|
|
||||||
JLT scalar_tail
|
|
||||||
|
|
||||||
VPXOR Y4, Y4, Y4
|
|
||||||
VPXOR Y5, Y5, Y5
|
|
||||||
VPXOR Y6, Y6, Y6
|
|
||||||
VPXOR Y7, Y7, Y7
|
|
||||||
|
|
||||||
CMPQ CX, $64
|
|
||||||
JLT loop32
|
|
||||||
|
|
||||||
loop64:
|
|
||||||
VPMOVZXDQ (SI), Y0
|
|
||||||
VPMOVZXDQ 16(SI), Y1
|
|
||||||
VPMOVZXDQ 32(SI), Y2
|
|
||||||
VPMOVZXDQ 48(SI), Y3
|
|
||||||
VPADDQ Y0, Y4, Y4
|
|
||||||
VPADDQ Y1, Y5, Y5
|
|
||||||
VPADDQ Y2, Y6, Y6
|
|
||||||
VPADDQ Y3, Y7, Y7
|
|
||||||
ADDQ $64, SI
|
|
||||||
SUBQ $64, CX
|
|
||||||
CMPQ CX, $64
|
|
||||||
JGE loop64
|
|
||||||
|
|
||||||
loop32:
|
|
||||||
CMPQ CX, $32
|
|
||||||
JLT reduce_vec
|
|
||||||
VPMOVZXDQ (SI), Y0
|
|
||||||
VPMOVZXDQ 16(SI), Y1
|
|
||||||
VPADDQ Y0, Y4, Y4
|
|
||||||
VPADDQ Y1, Y5, Y5
|
|
||||||
ADDQ $32, SI
|
|
||||||
SUBQ $32, CX
|
|
||||||
JMP loop32
|
|
||||||
|
|
||||||
reduce_vec:
|
|
||||||
// Combine the four ymm accumulators into Y4.
|
|
||||||
VPADDQ Y5, Y4, Y4
|
|
||||||
VPADDQ Y7, Y6, Y6
|
|
||||||
VPADDQ Y6, Y4, Y4
|
|
||||||
|
|
||||||
// Horizontally reduce Y4's four uint64 lanes to a single scalar.
|
|
||||||
VEXTRACTI128 $1, Y4, X5
|
|
||||||
VPADDQ X5, X4, X4
|
|
||||||
VPSHUFD $0x4e, X4, X5
|
|
||||||
VPADDQ X5, X4, X4
|
|
||||||
VMOVQ X4, R8
|
|
||||||
VZEROUPPER
|
|
||||||
|
|
||||||
ADDQ R8, AX
|
|
||||||
ADCQ $0, AX
|
|
||||||
|
|
||||||
scalar_tail:
|
|
||||||
// Handle remaining 0..63 bytes (or the entire buffer if it was < 32).
|
|
||||||
CMPQ CX, $8
|
|
||||||
JLT tail4
|
|
||||||
|
|
||||||
loop8:
|
|
||||||
ADDQ (SI), AX
|
|
||||||
ADCQ $0, AX
|
|
||||||
ADDQ $8, SI
|
|
||||||
SUBQ $8, CX
|
|
||||||
CMPQ CX, $8
|
|
||||||
JGE loop8
|
|
||||||
|
|
||||||
tail4:
|
|
||||||
CMPQ CX, $4
|
|
||||||
JLT tail2
|
|
||||||
MOVL (SI), R8
|
|
||||||
ADDQ R8, AX
|
|
||||||
ADCQ $0, AX
|
|
||||||
ADDQ $4, SI
|
|
||||||
SUBQ $4, CX
|
|
||||||
|
|
||||||
tail2:
|
|
||||||
CMPQ CX, $2
|
|
||||||
JLT tail1
|
|
||||||
MOVWQZX (SI), R8
|
|
||||||
ADDQ R8, AX
|
|
||||||
ADCQ $0, AX
|
|
||||||
ADDQ $2, SI
|
|
||||||
SUBQ $2, CX
|
|
||||||
|
|
||||||
tail1:
|
|
||||||
TESTQ CX, CX
|
|
||||||
JZ fold
|
|
||||||
MOVBQZX (SI), R8
|
|
||||||
ADDQ R8, AX
|
|
||||||
ADCQ $0, AX
|
|
||||||
|
|
||||||
fold:
|
|
||||||
// Fold the 64-bit accumulator to 16 bits via four rounds, mirroring
|
|
||||||
// gvisor's reduce(). Each pair (split, add) halves the live width;
|
|
||||||
// the truncation steps absorb the single bit that may be left over
|
|
||||||
// after each add so the next round's bound holds.
|
|
||||||
|
|
||||||
// 64 → 33 bits.
|
|
||||||
MOVQ AX, R8
|
|
||||||
SHRQ $32, R8
|
|
||||||
MOVL AX, AX
|
|
||||||
ADDQ R8, AX
|
|
||||||
|
|
||||||
// 33 → 32 bits. AX += (AX>>32); truncate to 32. AX is now ≤ 0xFFFF_FFFF.
|
|
||||||
MOVQ AX, R8
|
|
||||||
SHRQ $32, R8
|
|
||||||
ADDQ R8, AX
|
|
||||||
MOVL AX, AX
|
|
||||||
|
|
||||||
// 32 → 17 bits.
|
|
||||||
MOVQ AX, R8
|
|
||||||
SHRQ $16, R8
|
|
||||||
MOVWQZX AX, AX
|
|
||||||
ADDQ R8, AX
|
|
||||||
|
|
||||||
// 17 → 16 bits. AX += (AX>>16); the trailing MOVW truncates bit 16.
|
|
||||||
MOVQ AX, R8
|
|
||||||
SHRQ $16, R8
|
|
||||||
ADDQ R8, AX
|
|
||||||
|
|
||||||
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
|
||||||
// to big-endian to match the gvisor API contract.
|
|
||||||
XCHGB AH, AL
|
|
||||||
|
|
||||||
MOVW AX, ret+32(FP)
|
|
||||||
RET
|
|
||||||
@@ -1,12 +0,0 @@
|
|||||||
package checksum
|
|
||||||
|
|
||||||
//go:noescape
|
|
||||||
func checksumNEON(buf []byte, initial uint16) uint16
|
|
||||||
|
|
||||||
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
|
||||||
// initial. It is a drop-in replacement for gvisor's checksum.Checksum
|
|
||||||
// that dispatches to a hand-written NEON routine. NEON is mandatory in
|
|
||||||
// armv8 so no feature check is needed.
|
|
||||||
func Checksum(buf []byte, initial uint16) uint16 {
|
|
||||||
return checksumNEON(buf, initial)
|
|
||||||
}
|
|
||||||
@@ -1,143 +0,0 @@
|
|||||||
#include "textflag.h"
|
|
||||||
|
|
||||||
// func checksumNEON(buf []byte, initial uint16) uint16
|
|
||||||
//
|
|
||||||
// Mirrors the algorithm in checksum_amd64.s: sum the buffer treating it as
|
|
||||||
// a stream of uint32s in machine (little-endian) byte order, accumulating
|
|
||||||
// into 64-bit lanes that have ample carry headroom; fold and byte-swap once
|
|
||||||
// at the very end to recover the on-wire (big-endian) result.
|
|
||||||
//
|
|
||||||
// Each loop iteration loads 64 bytes via VLD1.P into V0..V3 (4 Q regs).
|
|
||||||
// VUADDW takes the low two uint32 lanes of a Q reg, zero-extends them to
|
|
||||||
// uint64, and adds them into a 2×uint64 accumulator; VUADDW2 does the same
|
|
||||||
// for the high two lanes. Four ymm-equivalent accumulators (V8..V11) get
|
|
||||||
// updated twice per iter to break the dep chain. Tail bytes go through a
|
|
||||||
// scalar ADCS chain seeded with the byte-swapped initial.
|
|
||||||
TEXT ·checksumNEON(SB), NOSPLIT, $0-34
|
|
||||||
MOVD buf_base+0(FP), R0
|
|
||||||
MOVD buf_len+8(FP), R1
|
|
||||||
MOVHU initial+24(FP), R2
|
|
||||||
|
|
||||||
// Pre-byteswap initial into the LE-summing space so it merges directly
|
|
||||||
// with the rest of the accumulator.
|
|
||||||
REV16W R2, R2
|
|
||||||
|
|
||||||
MOVD ZR, R3 // scalar accumulator
|
|
||||||
|
|
||||||
CMP $32, R1
|
|
||||||
BLT scalar_tail
|
|
||||||
|
|
||||||
VEOR V8.B16, V8.B16, V8.B16
|
|
||||||
VEOR V9.B16, V9.B16, V9.B16
|
|
||||||
VEOR V10.B16, V10.B16, V10.B16
|
|
||||||
VEOR V11.B16, V11.B16, V11.B16
|
|
||||||
|
|
||||||
CMP $64, R1
|
|
||||||
BLT loop16_init
|
|
||||||
|
|
||||||
loop64:
|
|
||||||
VLD1.P 64(R0), [V0.B16, V1.B16, V2.B16, V3.B16]
|
|
||||||
VUADDW V0.S2, V8.D2, V8.D2
|
|
||||||
VUADDW2 V0.S4, V9.D2, V9.D2
|
|
||||||
VUADDW V1.S2, V10.D2, V10.D2
|
|
||||||
VUADDW2 V1.S4, V11.D2, V11.D2
|
|
||||||
VUADDW V2.S2, V8.D2, V8.D2
|
|
||||||
VUADDW2 V2.S4, V9.D2, V9.D2
|
|
||||||
VUADDW V3.S2, V10.D2, V10.D2
|
|
||||||
VUADDW2 V3.S4, V11.D2, V11.D2
|
|
||||||
SUB $64, R1, R1
|
|
||||||
CMP $64, R1
|
|
||||||
BGE loop64
|
|
||||||
|
|
||||||
loop16_init:
|
|
||||||
CMP $16, R1
|
|
||||||
BLT reduce_vec
|
|
||||||
|
|
||||||
loop16:
|
|
||||||
VLD1.P 16(R0), [V0.B16]
|
|
||||||
VUADDW V0.S2, V8.D2, V8.D2
|
|
||||||
VUADDW2 V0.S4, V9.D2, V9.D2
|
|
||||||
SUB $16, R1, R1
|
|
||||||
CMP $16, R1
|
|
||||||
BGE loop16
|
|
||||||
|
|
||||||
reduce_vec:
|
|
||||||
// Combine the four accumulators into V8.
|
|
||||||
VADD V9.D2, V8.D2, V8.D2
|
|
||||||
VADD V11.D2, V10.D2, V10.D2
|
|
||||||
VADD V10.D2, V8.D2, V8.D2
|
|
||||||
|
|
||||||
// Horizontal-add the two lanes of V8.D2 into a single uint64.
|
|
||||||
VADDP V8.D2, V8.D2, V8.D2
|
|
||||||
VMOV V8.D[0], R8
|
|
||||||
|
|
||||||
ADDS R8, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
|
|
||||||
scalar_tail:
|
|
||||||
CMP $8, R1
|
|
||||||
BLT tail4
|
|
||||||
|
|
||||||
loop8:
|
|
||||||
MOVD.P 8(R0), R8
|
|
||||||
ADDS R8, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
SUB $8, R1, R1
|
|
||||||
CMP $8, R1
|
|
||||||
BGE loop8
|
|
||||||
|
|
||||||
tail4:
|
|
||||||
CMP $4, R1
|
|
||||||
BLT tail2
|
|
||||||
MOVWU.P 4(R0), R8
|
|
||||||
ADDS R8, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
SUB $4, R1, R1
|
|
||||||
|
|
||||||
tail2:
|
|
||||||
CMP $2, R1
|
|
||||||
BLT tail1
|
|
||||||
MOVHU.P 2(R0), R8
|
|
||||||
ADDS R8, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
SUB $2, R1, R1
|
|
||||||
|
|
||||||
tail1:
|
|
||||||
CBZ R1, fold
|
|
||||||
MOVBU (R0), R8
|
|
||||||
ADDS R8, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
|
|
||||||
fold:
|
|
||||||
// Merge the byte-swapped initial into our LE-form accumulator.
|
|
||||||
ADDS R2, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
|
|
||||||
// 64 → 33 bits.
|
|
||||||
LSR $32, R3, R8
|
|
||||||
AND $0xffffffff, R3, R3
|
|
||||||
ADD R8, R3, R3
|
|
||||||
|
|
||||||
// 33 → 32 (truncate after adding bit 32 back).
|
|
||||||
LSR $32, R3, R8
|
|
||||||
ADD R8, R3, R3
|
|
||||||
AND $0xffffffff, R3, R3
|
|
||||||
|
|
||||||
// 32 → 17.
|
|
||||||
LSR $16, R3, R8
|
|
||||||
AND $0xffff, R3, R3
|
|
||||||
ADD R8, R3, R3
|
|
||||||
|
|
||||||
// 17 → 16 (truncation absorbs bit 16 below).
|
|
||||||
LSR $16, R3, R8
|
|
||||||
ADD R8, R3, R3
|
|
||||||
|
|
||||||
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
|
||||||
// to big-endian to match the gvisor API contract. REV16W swaps bytes
|
|
||||||
// within each 16-bit halfword of the low 32 bits, so it acts as a
|
|
||||||
// 16-bit byte-swap on the live low 16.
|
|
||||||
REV16W R3, R3
|
|
||||||
AND $0xffff, R3, R3
|
|
||||||
|
|
||||||
MOVH R3, ret+32(FP)
|
|
||||||
RET
|
|
||||||
@@ -1,10 +0,0 @@
|
|||||||
//go:build !amd64 && !arm64
|
|
||||||
|
|
||||||
package checksum
|
|
||||||
|
|
||||||
import gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
|
||||||
|
|
||||||
// Checksum delegates to gvisor on architectures without a hand-written body.
|
|
||||||
func Checksum(buf []byte, initial uint16) uint16 {
|
|
||||||
return gvisorchecksum.Checksum(buf, initial)
|
|
||||||
}
|
|
||||||
@@ -1,232 +0,0 @@
|
|||||||
package checksum
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"math/rand/v2"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
|
||||||
)
|
|
||||||
|
|
||||||
// archImpl names one checksum function under test. The per-arch
|
|
||||||
// export_*_test.go files enumerate the hand-written implementations so the
|
|
||||||
// suite compares each one against gvisor directly, regardless of which one
|
|
||||||
// the public Checksum dispatches to on the running CPU. Testing only the
|
|
||||||
// dispatcher was tautological wherever it resolved to the gvisor fallback
|
|
||||||
// (non-AVX2 amd64, fallback architectures) — gvisor compared with itself,
|
|
||||||
// assembly untested, suite green.
|
|
||||||
type archImpl struct {
|
|
||||||
name string
|
|
||||||
fn func([]byte, uint16) uint16
|
|
||||||
available bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// implsUnderTest is the public dispatcher plus every arch implementation.
|
|
||||||
func implsUnderTest() []archImpl {
|
|
||||||
return append([]archImpl{{name: "dispatch", fn: Checksum, available: true}}, archImpls...)
|
|
||||||
}
|
|
||||||
|
|
||||||
// requireAvailable skips loudly when the running CPU can't execute an
|
|
||||||
// implementation — visible in test output, unlike the old silent tautology.
|
|
||||||
func requireAvailable(t *testing.T, impl archImpl) {
|
|
||||||
t.Helper()
|
|
||||||
if !impl.available {
|
|
||||||
t.Skipf("%s not supported on this CPU; its assembly is NOT tested in this run", impl.name)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
|
|
||||||
// seeds and a handful of starting alignments, asserting that each local
|
|
||||||
// implementation matches gvisor's reference bit-for-bit.
|
|
||||||
func TestChecksumMatchesGvisor(t *testing.T) {
|
|
||||||
for _, impl := range implsUnderTest() {
|
|
||||||
t.Run(impl.name, func(t *testing.T) {
|
|
||||||
requireAvailable(t, impl)
|
|
||||||
rng := rand.New(rand.NewPCG(1, 2))
|
|
||||||
const padFront = 16
|
|
||||||
|
|
||||||
// Random pool large enough for the longest case + alignment slop.
|
|
||||||
pool := make([]byte, 4096+padFront)
|
|
||||||
for i := range pool {
|
|
||||||
pool[i] = byte(rng.Uint32())
|
|
||||||
}
|
|
||||||
|
|
||||||
seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc}
|
|
||||||
offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16}
|
|
||||||
|
|
||||||
for length := 0; length <= 4096; length++ {
|
|
||||||
for _, seed := range seeds {
|
|
||||||
for _, off := range offsets {
|
|
||||||
if off+length > len(pool) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
buf := pool[off : off+length]
|
|
||||||
want := gvisorchecksum.Checksum(buf, seed)
|
|
||||||
got := impl.fn(buf, seed)
|
|
||||||
if got != want {
|
|
||||||
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
|
|
||||||
length, off, seed, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestChecksumPatternedBuffers exercises specific byte patterns that have
|
|
||||||
// historically tripped up checksum implementations: all-zero, all-0xff,
|
|
||||||
// alternating, and ascending sequences.
|
|
||||||
func TestChecksumPatternedBuffers(t *testing.T) {
|
|
||||||
for _, impl := range implsUnderTest() {
|
|
||||||
t.Run(impl.name, func(t *testing.T) {
|
|
||||||
requireAvailable(t, impl)
|
|
||||||
for length := 0; length <= 256; length++ {
|
|
||||||
patterns := map[string][]byte{
|
|
||||||
"zeros": make([]byte, length),
|
|
||||||
"ones": bytes(length, 0xff),
|
|
||||||
"alternating": pattern(length, []byte{0xa5, 0x5a}),
|
|
||||||
"ascending": ascending(length),
|
|
||||||
}
|
|
||||||
for name, buf := range patterns {
|
|
||||||
for _, seed := range []uint16{0, 0xffff, 0x8000} {
|
|
||||||
want := gvisorchecksum.Checksum(buf, seed)
|
|
||||||
got := impl.fn(buf, seed)
|
|
||||||
if got != want {
|
|
||||||
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
|
|
||||||
name, length, seed, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func bytes(n int, v byte) []byte {
|
|
||||||
b := make([]byte, n)
|
|
||||||
for i := range b {
|
|
||||||
b[i] = v
|
|
||||||
}
|
|
||||||
return b
|
|
||||||
}
|
|
||||||
|
|
||||||
func pattern(n int, p []byte) []byte {
|
|
||||||
b := make([]byte, n)
|
|
||||||
for i := range b {
|
|
||||||
b[i] = p[i%len(p)]
|
|
||||||
}
|
|
||||||
return b
|
|
||||||
}
|
|
||||||
|
|
||||||
func ascending(n int) []byte {
|
|
||||||
b := make([]byte, n)
|
|
||||||
for i := range b {
|
|
||||||
b[i] = byte(i)
|
|
||||||
}
|
|
||||||
return b
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestChecksumTailPaths targets every combination of (SIMD body iterations,
|
|
||||||
// trailing tail bytes) the asm handlers walk through. The tail handlers
|
|
||||||
// peel off 8 → 4 → 2 → 1 byte chunks in turn; this test exercises each by
|
|
||||||
// constructing lengths of the form 64*k + tail for tail ∈ [0, 63] and a
|
|
||||||
// representative spread of k values, including k=0 (no main loop, all tail)
|
|
||||||
// and k=1 (one main loop iter, then tail). It's explicit coverage for
|
|
||||||
// payload sizes that are odd, not divisible by 4, by 8, or by 32.
|
|
||||||
func TestChecksumTailPaths(t *testing.T) {
|
|
||||||
for _, impl := range implsUnderTest() {
|
|
||||||
t.Run(impl.name, func(t *testing.T) {
|
|
||||||
requireAvailable(t, impl)
|
|
||||||
rng := rand.New(rand.NewPCG(42, 17))
|
|
||||||
const padFront = 16
|
|
||||||
const maxK = 8
|
|
||||||
|
|
||||||
pool := make([]byte, 64*maxK+padFront+64)
|
|
||||||
for i := range pool {
|
|
||||||
pool[i] = byte(rng.Uint32())
|
|
||||||
}
|
|
||||||
|
|
||||||
seeds := []uint16{0, 0xffff, 0xabcd}
|
|
||||||
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
|
|
||||||
|
|
||||||
for k := 0; k <= maxK; k++ {
|
|
||||||
for tail := 0; tail < 64; tail++ {
|
|
||||||
length := 64*k + tail
|
|
||||||
for _, seed := range seeds {
|
|
||||||
for _, off := range offsets {
|
|
||||||
if off+length > len(pool) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
buf := pool[off : off+length]
|
|
||||||
want := gvisorchecksum.Checksum(buf, seed)
|
|
||||||
got := impl.fn(buf, seed)
|
|
||||||
if got != want {
|
|
||||||
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
|
|
||||||
k, tail, length, off, seed, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkChecksumTailSizes covers payload sizes that aren't clean multiples
|
|
||||||
// of the SIMD body's 32-byte (amd64) or 16-byte (arm64) chunks, so the tail
|
|
||||||
// handler is meaningfully on the hot path. Sizes are picked to either exercise
|
|
||||||
// every tail branch (tiny lengths) or sit slightly off realistic packet
|
|
||||||
// boundaries (e.g. 1499 = MTU − 1).
|
|
||||||
func BenchmarkChecksumTailSizes(b *testing.B) {
|
|
||||||
sizes := []int{
|
|
||||||
1, 3, 7, 15, 31, // sub-SIMD; entire work is scalar tail
|
|
||||||
33, 35, 47, 63, // one loop32 + assorted tails
|
|
||||||
65, 95, 127, // one loop64 + assorted tails
|
|
||||||
1447, 1471, 1499, 1501, // around MTU
|
|
||||||
8191, 8193, // around USO
|
|
||||||
65531, 65533, // near the kernel max
|
|
||||||
}
|
|
||||||
for _, size := range sizes {
|
|
||||||
buf := make([]byte, size)
|
|
||||||
for i := range buf {
|
|
||||||
buf[i] = byte(i)
|
|
||||||
}
|
|
||||||
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
|
|
||||||
b.SetBytes(int64(size))
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
_ = Checksum(buf, 0)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
|
|
||||||
b.SetBytes(int64(size))
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
_ = gvisorchecksum.Checksum(buf, 0)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkChecksum compares the local Checksum to gvisor's at sizes that
|
|
||||||
// match real traffic: a TCP/IP header (60), a typical MSS (1448), a typical
|
|
||||||
// USO size (8192), and the kernel's max GSO superpacket (65535).
|
|
||||||
func BenchmarkChecksum(b *testing.B) {
|
|
||||||
for _, size := range []int{60, 1448, 8192, 65535} {
|
|
||||||
buf := make([]byte, size)
|
|
||||||
for i := range buf {
|
|
||||||
buf[i] = byte(i)
|
|
||||||
}
|
|
||||||
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
|
|
||||||
b.SetBytes(int64(size))
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
_ = Checksum(buf, 0)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
|
|
||||||
b.SetBytes(int64(size))
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
_ = gvisorchecksum.Checksum(buf, 0)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,11 +0,0 @@
|
|||||||
package checksum
|
|
||||||
|
|
||||||
// archImpls exposes every hand-written implementation on this architecture
|
|
||||||
// so the tests exercise them directly, independent of what the public
|
|
||||||
// Checksum dispatches to on the running CPU. Without this, running the
|
|
||||||
// suite on a non-AVX2 machine compared gvisor against itself and left the
|
|
||||||
// assembly untested — silently. available=false makes the test skip loudly
|
|
||||||
// instead.
|
|
||||||
var archImpls = []archImpl{
|
|
||||||
{name: "avx2", fn: checksumAVX2, available: hasAVX2},
|
|
||||||
}
|
|
||||||
@@ -1,8 +0,0 @@
|
|||||||
package checksum
|
|
||||||
|
|
||||||
// archImpls exposes every hand-written implementation on this architecture
|
|
||||||
// for direct testing; see export_amd64_test.go for the rationale. NEON is
|
|
||||||
// mandatory in armv8, so it is always available.
|
|
||||||
var archImpls = []archImpl{
|
|
||||||
{name: "neon", fn: checksumNEON, available: true},
|
|
||||||
}
|
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
//go:build !amd64 && !arm64
|
|
||||||
|
|
||||||
package checksum
|
|
||||||
|
|
||||||
// No hand-written implementations on this architecture; the dispatcher is
|
|
||||||
// pure gvisor and there is nothing separate to test.
|
|
||||||
var archImpls []archImpl
|
|
||||||
+3
-13
@@ -4,25 +4,15 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
// defaultBatchBufSize is the per-Queue scratch size for Read on backends
|
|
||||||
// that don't do TSO segmentation. 65535 covers any single IP packet.
|
|
||||||
const defaultBatchBufSize = 65535
|
|
||||||
|
|
||||||
type Device interface {
|
type Device interface {
|
||||||
io.Closer
|
io.ReadWriteCloser
|
||||||
Activate() error
|
Activate() error
|
||||||
Networks() []netip.Prefix
|
Networks() []netip.Prefix
|
||||||
Name() string
|
Name() string
|
||||||
RoutesFor(netip.Addr) routing.Gateways
|
RoutesFor(netip.Addr) routing.Gateways
|
||||||
// Queues returns the device's packet queues, opening additional ones as
|
SupportsMultiqueue() bool
|
||||||
// needed until there are n. Platforms without multiqueue support return
|
NewMultiQueueReader() (io.ReadWriteCloser, error)
|
||||||
// their single queue regardless of n, so callers must size reader loops
|
|
||||||
// to len(result), not n; implementations never return more than n. An
|
|
||||||
// error means a queue that should have opened could not; the caller owns
|
|
||||||
// cleanup via Close. Called once, during interface activation.
|
|
||||||
Queues(n int) ([]tio.Queue, error)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,9 +3,10 @@
|
|||||||
package overlaytest
|
package overlaytest
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -30,16 +31,20 @@ func (NoopTun) Name() string {
|
|||||||
return "noop"
|
return "noop"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Read() ([]tio.Packet, error) {
|
func (NoopTun) Read([]byte) (int, error) {
|
||||||
return nil, nil
|
return 0, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Write([]byte) (int, error) {
|
func (NoopTun) Write([]byte) (int, error) {
|
||||||
return 0, nil
|
return 0, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Queues(int) ([]tio.Queue, error) {
|
func (NoopTun) SupportsMultiqueue() bool {
|
||||||
return []tio.Queue{NoopTun{}}, nil
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func (NoopTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
|
return nil, errors.New("unsupported")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Close() error {
|
func (NoopTun) Close() error {
|
||||||
|
|||||||
@@ -1,44 +0,0 @@
|
|||||||
//go:build linux && !android
|
|
||||||
// +build linux,!android
|
|
||||||
|
|
||||||
package tio
|
|
||||||
|
|
||||||
import (
|
|
||||||
"os"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
)
|
|
||||||
|
|
||||||
// blockOn parks the calling goroutine until fd is ready or shutdownFd signals teardown.
|
|
||||||
// (events is POLLIN for reads, POLLOUT for writes)
|
|
||||||
// It builds the pollfd array on the stack every call, so concurrent callers on the same Queue never share Revents storage.
|
|
||||||
//
|
|
||||||
// Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd)
|
|
||||||
// or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR).
|
|
||||||
func blockOn(fd, shutdownFd int32, events int16) error {
|
|
||||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
|
||||||
pfds := [2]unix.PollFd{
|
|
||||||
{Fd: fd, Events: events},
|
|
||||||
{Fd: shutdownFd, Events: unix.POLLIN},
|
|
||||||
}
|
|
||||||
var err error
|
|
||||||
for {
|
|
||||||
_, err = unix.Poll(pfds[:], -1)
|
|
||||||
if err != unix.EINTR {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
tunEvents := pfds[0].Revents
|
|
||||||
shutdownEvents := pfds[1].Revents
|
|
||||||
// Check err before trusting the potentially bogus bits we just got.
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
|
||||||
return os.ErrClosed
|
|
||||||
}
|
|
||||||
if tunEvents&problemFlags != 0 {
|
|
||||||
return os.ErrClosed
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user