mirror of
https://github.com/slackhq/nebula.git
synced 2026-09-30 09:46:37 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5f20a6adcf | ||
|
|
e50f8128f4 | ||
|
|
dd8f660c0a | ||
|
|
ec3304e3a9 | ||
|
|
aaa2ff7fff | ||
|
|
657f6ad044 | ||
|
|
19e04db115 | ||
|
|
9c5d701648 | ||
|
|
edc3c5e018 | ||
|
|
b8b159a486 | ||
|
|
49e35d1283 | ||
|
|
6fcb926334 | ||
|
|
28d82f7b8a | ||
|
|
f15d10fc54 | ||
|
|
bfe7790848 | ||
|
|
e7b4e094b9 | ||
|
|
6740518403 | ||
|
|
5cbf029d06 | ||
|
|
368855c2e8 | ||
|
|
6d124d0441 | ||
|
|
72bf111209 | ||
|
|
1617897043 | ||
|
|
f8775bb6ca | ||
|
|
15f0f0d5d0 | ||
|
|
7902ce674e | ||
|
|
c2fbe215e6 | ||
|
|
94ac6db4ca | ||
|
|
a60350e34e | ||
|
|
58f3b6fda7 | ||
|
|
a99699e370 | ||
|
|
3615a79b8b | ||
|
|
147c202c27 | ||
|
|
e290a6892f |
@@ -43,8 +43,15 @@ runs:
|
||||
with:
|
||||
role-to-assume: ${{ inputs.role }}
|
||||
aws-region: ${{ inputs.region }}
|
||||
# Default is 12 retries to ride out IAM trust-policy propagation; once
|
||||
# the role is stable we want a real misconfiguration to fail fast.
|
||||
# An STS secret key with special characters does not survive the
|
||||
# pwsh -> make -> MSYS sh -> aws.exe chain, and SigV4 then signs with a
|
||||
# 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
|
||||
|
||||
- name: Sign .exe files
|
||||
|
||||
@@ -12,9 +12,9 @@ jobs:
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -38,9 +38,9 @@ jobs:
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -73,27 +73,81 @@ jobs:
|
||||
build-darwin:
|
||||
name: Build Universal Darwin
|
||||
env:
|
||||
HAS_SIGNING_CREDS: ${{ secrets.AC_USERNAME != '' }}
|
||||
HAS_SIGNING_CREDS: ${{ secrets.APPLE_SIGNING_ROLE_ARN != '' }}
|
||||
runs-on: macos-latest
|
||||
permissions:
|
||||
id-token: write
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
# GitHub holds ARNs, not credentials, and ARNs outlive a rotation
|
||||
- name: Configure AWS credentials
|
||||
if: env.HAS_SIGNING_CREDS == 'true'
|
||||
uses: aws-actions/configure-aws-credentials@v6
|
||||
with:
|
||||
role-to-assume: ${{ secrets.APPLE_SIGNING_ROLE_ARN }}
|
||||
aws-region: us-east-2
|
||||
|
||||
# parse-json-secrets unpacks into SIGNING_* and ASC_*, masked on the way in
|
||||
- name: Fetch signing credentials
|
||||
if: env.HAS_SIGNING_CREDS == 'true'
|
||||
uses: aws-actions/aws-secretsmanager-get-secrets@v3
|
||||
with:
|
||||
parse-json-secrets: true
|
||||
secret-ids: |
|
||||
SIGNING,${{ secrets.APPLE_SIGNING_DEVELOPER_ID_ARN }}
|
||||
ASC,${{ secrets.APPLE_NOTARY_KEY_ARN }}
|
||||
|
||||
- name: Import certificates
|
||||
if: env.HAS_SIGNING_CREDS == 'true'
|
||||
uses: Apple-Actions/import-codesign-certs@v7
|
||||
with:
|
||||
p12-file-base64: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_P12_BASE64 }}
|
||||
p12-password: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_PASSWORD }}
|
||||
p12-file-base64: ${{ env.SIGNING_P12_BASE64 }}
|
||||
p12-password: ${{ env.SIGNING_PASSWORD }}
|
||||
|
||||
# The action imports but does not check the chain validates, which is how a p12
|
||||
# missing its intermediate reaches a failing codesign
|
||||
- name: Check the identity is usable
|
||||
if: env.HAS_SIGNING_CREDS == 'true'
|
||||
run: |
|
||||
: "${SIGNING_IDENTITY_SHA1:?empty, so the secret has no identity_sha1}"
|
||||
identities=$(security find-identity -v -p codesigning signing_temp.keychain)
|
||||
case "$identities" in
|
||||
*"$SIGNING_IDENTITY_SHA1"*) ;;
|
||||
*) printf '%s\n' "$identities" >&2; exit 1 ;;
|
||||
esac
|
||||
|
||||
# notarytool wants the key as a file
|
||||
- name: Write the App Store Connect key
|
||||
if: env.HAS_SIGNING_CREDS == 'true'
|
||||
run: |
|
||||
mkdir -p ~/private_keys
|
||||
chmod 700 ~/private_keys
|
||||
key_path="$HOME/private_keys/AuthKey_${ASC_KEY_ID}.p8"
|
||||
(umask 077; printf '%s\n' "$ASC_PRIVATE_KEY" > "$key_path")
|
||||
echo "ASC_P8=$key_path" >> "$GITHUB_ENV"
|
||||
|
||||
- name: Drop the credentials from the environment
|
||||
if: env.HAS_SIGNING_CREDS == 'true'
|
||||
run: |
|
||||
# The action's own inventory, so a new field in a secret is covered
|
||||
python3 -c '
|
||||
import json, os
|
||||
raw = os.environ.get("SECRETS_LIST_CLEAN_UP")
|
||||
if raw is None and os.environ.get("SIGNING_P12_BASE64"):
|
||||
raise SystemExit("SECRETS_LIST_CLEAN_UP is gone, fetched secrets are not being scrubbed")
|
||||
keep = {"SIGNING_IDENTITY_SHA1", "ASC_KEY_ID", "ASC_ISSUER_ID"}
|
||||
names = [n for n in json.loads(raw or "[]") if n not in keep]
|
||||
print("\n".join(f"{n}=" for n in dict.fromkeys(names)))
|
||||
' >> "$GITHUB_ENV"
|
||||
|
||||
- name: Build, sign, and notarize
|
||||
env:
|
||||
AC_USERNAME: ${{ secrets.AC_USERNAME }}
|
||||
AC_PASSWORD: ${{ secrets.AC_PASSWORD }}
|
||||
run: |
|
||||
rm -rf release
|
||||
mkdir release
|
||||
@@ -102,17 +156,34 @@ jobs:
|
||||
lipo -create -output ./release/nebula ./build/darwin-amd64/nebula ./build/darwin-arm64/nebula
|
||||
lipo -create -output ./release/nebula-cert ./build/darwin-amd64/nebula-cert ./build/darwin-arm64/nebula-cert
|
||||
|
||||
if [ -n "$AC_USERNAME" ]; then
|
||||
codesign -s "10BC1FDDEB6CE753550156C0669109FAC49E4D1E" -f -v --timestamp --options=runtime -i "net.defined.nebula" ./release/nebula
|
||||
codesign -s "10BC1FDDEB6CE753550156C0669109FAC49E4D1E" -f -v --timestamp --options=runtime -i "net.defined.nebula-cert" ./release/nebula-cert
|
||||
# Unset in a fork, which has no credentials to sign with
|
||||
if [ -n "$SIGNING_IDENTITY_SHA1" ]; then
|
||||
codesign -s "$SIGNING_IDENTITY_SHA1" -f -v --timestamp --options=runtime -i "net.defined.nebula" ./release/nebula
|
||||
codesign -s "$SIGNING_IDENTITY_SHA1" -f -v --timestamp --options=runtime -i "net.defined.nebula-cert" ./release/nebula-cert
|
||||
fi
|
||||
|
||||
zip -j release/nebula-darwin.zip release/nebula-cert release/nebula
|
||||
|
||||
if [ -n "$AC_USERNAME" ]; then
|
||||
xcrun notarytool submit ./release/nebula-darwin.zip --team-id "576H3XS7FP" --apple-id "$AC_USERNAME" --password "$AC_PASSWORD" --wait
|
||||
if [ -n "$ASC_P8" ]; then
|
||||
xcrun notarytool submit ./release/nebula-darwin.zip --key "$ASC_P8" --key-id "$ASC_KEY_ID" --issuer "$ASC_ISSUER_ID" --wait
|
||||
fi
|
||||
|
||||
- name: Drop the signing key
|
||||
if: always() && env.HAS_SIGNING_CREDS == 'true'
|
||||
run: |
|
||||
# Locked, not deleted: import-codesign-certs deletes it in its own post
|
||||
# step and fails the job if it is already gone. Locked is unusable.
|
||||
security lock-keychain signing_temp.keychain || true
|
||||
rm -f "$ASC_P8"
|
||||
# Nothing later in this job needs AWS
|
||||
python3 -c '
|
||||
import json, os
|
||||
names = json.loads(os.environ.get("SECRETS_LIST_CLEAN_UP") or "[]")
|
||||
names += ["ASC_P8", "SIGNING_IDENTITY_SHA1", "ASC_KEY_ID", "ASC_ISSUER_ID",
|
||||
"AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_SESSION_TOKEN"]
|
||||
print("\n".join(f"{n}=" for n in dict.fromkeys(names)))
|
||||
' >> "$GITHUB_ENV"
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v7
|
||||
with:
|
||||
|
||||
@@ -32,9 +32,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: add hashicorp source
|
||||
@@ -64,9 +64,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: add hashicorp source
|
||||
@@ -90,9 +90,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
# WSL2 + Ubuntu so the smoke can run a real linux peer with its own
|
||||
|
||||
+34
-33
@@ -20,44 +20,45 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
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
|
||||
with:
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: build
|
||||
run: make bin-docker CGO_ENABLED=1 BUILD_ARGS=-race
|
||||
run: make bin
|
||||
|
||||
- name: setup docker image
|
||||
- name: run smoke-self
|
||||
working-directory: ./.github/workflows/smoke
|
||||
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
|
||||
run: ./smoke-self.sh
|
||||
|
||||
timeout-minutes: 10
|
||||
|
||||
Executable
+130
@@ -0,0 +1,130 @@
|
||||
#!/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
|
||||
+15
-10
@@ -20,9 +20,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Install goimports
|
||||
@@ -42,7 +42,7 @@ jobs:
|
||||
- name: golangci-lint
|
||||
uses: golangci/golangci-lint-action@v9
|
||||
with:
|
||||
version: v2.5
|
||||
version: v2.12
|
||||
|
||||
test:
|
||||
name: Test ${{ matrix.name }}
|
||||
@@ -58,9 +58,14 @@ jobs:
|
||||
e2e-cmd: make e2evv
|
||||
- name: linux-boringcrypto
|
||||
os: ubuntu-latest
|
||||
build-cmd: make bin-boringcrypto
|
||||
test-cmd: make test-boringcrypto
|
||||
e2e-cmd: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0"
|
||||
build-cmd: make boringcrypto
|
||||
test-cmd: make boringcrypto test
|
||||
e2e-cmd: make boringcrypto e2evv
|
||||
- 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
|
||||
os: ubuntu-latest
|
||||
build-cmd: make bin-pkcs11
|
||||
@@ -80,9 +85,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -125,9 +130,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build ${{ matrix.name }}
|
||||
|
||||
@@ -7,6 +7,102 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
### 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)
|
||||
|
||||
### 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)
|
||||
|
||||
## [1.11.0] - 2026-07-23
|
||||
|
||||
See the [v1.11.0](https://github.com/slackhq/nebula/milestone/25?closed=1) milestone for a complete list of changes.
|
||||
|
||||
### Breaking
|
||||
|
||||
- Logging has switched from logrus to Go's structured `slog`. Log output changes: levels are upper case
|
||||
(`level=INFO`), trace prints as `level=DEBUG-4`, timestamps are always RFC3339Nano and `logging.timestamp_format`
|
||||
is ignored, and some messages were reworded. Review any log parsing before upgrading. This is also an API break
|
||||
for embedders, as constructors now take a `*slog.Logger`. (#1672, #1734, #1621)
|
||||
- `firewall.inbound_action` and `firewall.outbound_action` (used to set reject vs. drop policy) were each being
|
||||
applied to the opposite direction, that is now corrected. This only affects how blocked packets are answered, not
|
||||
which packets the firewall allows or denies. If you set either of these you are getting the behavior of the other
|
||||
one today and likely want to swap them before upgrading. (#1798)
|
||||
- On Windows, Nebula now installs WFP PERMIT filters for the nebula adapter and the listener port by default. WFP
|
||||
sits below Windows Defender Firewall, so any WDF inbound rules you rely on for either will no longer apply. Set
|
||||
`tun.windows_bypass_wdf` and `listen.windows_bypass_wdf` to false to leave WDF in charge. (#1710)
|
||||
- On Windows, the nebula device is now set to the `private` network category instead of whatever Windows decided,
|
||||
which is usually `Public`. This makes the host firewall less restrictive on the overlay. Set
|
||||
`tun.network_category` to `unset` to keep the old behavior. (#1710)
|
||||
- Reject packets for non-TCP now use ICMP code 13, communication administratively prohibited, instead of code 3,
|
||||
port unreachable. Anything keying off the old code needs updating. (#1766, #1768)
|
||||
- The SSH debug server's profiling commands are now confined to `sshd.sandbox_dir`, which defaults to
|
||||
`$TMP/nebula-debug`. Relative paths resolve inside it and absolute paths outside it are rejected, so anything
|
||||
scripting `start-cpu-profile`, `save-heap-profile`, or `save-mutex-profile` with a path elsewhere needs the
|
||||
directory set. The directory is not created for you. (#1622)
|
||||
|
||||
### Added
|
||||
|
||||
- Sign the Windows release binaries. (#1718)
|
||||
- Generate IPv6 reject packets, matching the existing IPv4 behavior. (#1766, #1767, #1768)
|
||||
- Accept `-` in `nebula-cert` to read from stdin or write to stdout. (#1714)
|
||||
- Search for both `config.yml` and `config.yaml` in service and command line modes. (#1717)
|
||||
- Add version labels to the Docker/OCI images. (#1772)
|
||||
- Rebind the listener and re-query lighthouses on macOS when the underlay network changes, so devices moving
|
||||
between wifi and wired or between networks recover without waiting for dead tunnel detection. Controlled by
|
||||
`listen.rebind_on_network_change` (default `true`, not reloadable). (#1816)
|
||||
|
||||
### Changed
|
||||
|
||||
- Reload the firewall when the unsafe networks in the certificate change. (#1719)
|
||||
- Reconfigure, start, and stop the stats listener on a config reload instead of requiring a restart. (#1670)
|
||||
- Update a static host's addresses when they change on reload. (#1713)
|
||||
- Don't require a port on ICMP firewall rules. (#1609)
|
||||
- Connection track ICMP traffic. (#1602)
|
||||
- Return `NODATA` instead of `NXDOMAIN` from the DNS server for a name that exists but has no record of the
|
||||
requested type, so clients that query `AAAA` first (busybox/Alpine) fall through to `A`. (#1668)
|
||||
- Record the local host's details in the DNS server. (#1716)
|
||||
- Install Windows unsafe routes as link routes. (#1709)
|
||||
- Reduce relay handshake log spam, and only log a handshake send error at error level when the remote list
|
||||
changes. (#1733, #1765, #1810)
|
||||
- Start, stop, and reload subsystems (DNS, stats, conntrack, ssh, punchy) cleanly without leaking goroutines. (#1640, #1654, #1661, #1667, #1669, #1708, #1806, #1815)
|
||||
- `Control` is now safe to stop and wait on from any lifecycle state, and a new `Control.Wait` blocks until nebula
|
||||
has fully stopped and returns the first fatal reader error. Failed starts release the udp sockets and tun fd
|
||||
instead of leaking them. (#1794)
|
||||
- Trigger an immediate lighthouse update when reconnecting to or adding a lighthouse instead of waiting for the next update tick. (#1645)
|
||||
- Bring the Darwin and OpenBSD tun implementations in line with the other BSDs. (#1703)
|
||||
- Update to build against go v1.26. (#1818)
|
||||
- Various dependency updates. (#1586, #1587, #1604, #1617, #1618, #1627, #1628, #1629, #1652, #1664, #1665, #1697, #1721, #1732, #1742, #1743, #1750, #1763, #1771, #1782, #1800, #1807)
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix a data race on a host's remote address that could send packets to the wrong address during a roam. (#1773)
|
||||
- Fix tunnels that could permanently escape connection manager monitoring. (#1752)
|
||||
- Fix a crash when reloading the SSH server's trusted keys. (#1787)
|
||||
- Fix hostmap corruption when a host has multiple overlay addresses. Each address now gets its own list instead of
|
||||
a single shared chain, which also fixes two latent bugs on the add and makePrimary paths. (#1788, #1790)
|
||||
- Apply `remote_allow_list` IPv4 rules to 4-in-6 mapped addresses. (#1786)
|
||||
- Don't panic in the DNS server on a short or empty query name. (#1635)
|
||||
- Advance the replay window on relayed packets so a relay drops replayed frames instead of re-forwarding them. (#1751)
|
||||
- Fix a race in relay state handling. (#1753)
|
||||
- Lock replay window updates so concurrent readers can't corrupt it. (#1802)
|
||||
- Reject malformed handshakes more reliably, including invalid ed25519 key lengths. (#1601, #1756)
|
||||
- Properly handle `closetunnel` packets. (#1638)
|
||||
- Fix an IPv6 extension-header length overflow that could make the firewall parse the wrong protocol and ports. (#1789)
|
||||
- Fix relay re-establishment when a handshake arrives over a relay entry that a one-sided teardown left
|
||||
`Disestablished`, which silently dropped every send until dead tunnel detection forced a re-handshake. (#1805)
|
||||
- Don't build new relay state on a tunnel that was just discarded. (#1796)
|
||||
- Don't delete the wrong pending hostinfo in the handshake manager. (#1811)
|
||||
- Don't call the packet reader after a UDP error on Darwin. (#1755)
|
||||
- Open the FreeBSD tun device non blocking. (#1666)
|
||||
|
||||
## [1.10.3] - 2026-02-06
|
||||
|
||||
### Security
|
||||
|
||||
@@ -72,6 +72,17 @@ ALL_CROSS_LINUX_ARM = linux-arm-5 linux-arm-6 linux-arm-7 linux-arm64
|
||||
ALL_CROSS_LINUX_MIPS = linux-mips linux-mipsle linux-mips64 linux-mips64le linux-mips-softfloat
|
||||
ALL_CROSS_LINUX_OTHER = linux-386 linux-ppc64le linux-riscv64 linux-loong64
|
||||
|
||||
# Based on section 2.2 of the Go Cryptographic Module CVMP Security Policy #5247
|
||||
ALL_FIPS140 = linux-amd64-fips140 \
|
||||
linux-arm64-fips140 \
|
||||
windows-amd64-fips140 \
|
||||
windows-arm64-fips140 \
|
||||
darwin-arm64-fips140 \
|
||||
freebsd-amd64-fips140 \
|
||||
linux-arm-7-fips140 \
|
||||
linux-mips64-fips140 \
|
||||
linux-ppc64le-fips140
|
||||
|
||||
e2e:
|
||||
$(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e
|
||||
|
||||
@@ -137,6 +148,8 @@ release-netbsd: $(ALL_NETBSD:%=build/nebula-%.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
|
||||
|
||||
bin-windows: build/windows-amd64/nebula.exe build/windows-amd64/nebula-cert.exe
|
||||
@@ -157,6 +170,9 @@ 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
|
||||
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: CGO_ENABLED = 1
|
||||
bin-pkcs11: bin
|
||||
@@ -166,12 +182,12 @@ debug: BUILD_ARGS += -tags debug
|
||||
debug: bin
|
||||
|
||||
bin:
|
||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
||||
$(GOENV) 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
|
||||
|
||||
install:
|
||||
go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ${NEBULA_CMD_PATH}
|
||||
go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ./cmd/nebula-cert
|
||||
$(GOENV) go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ${NEBULA_CMD_PATH}
|
||||
$(GOENV) go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ./cmd/nebula-cert
|
||||
|
||||
build/linux-arm-%: GOENV += GOARM=$(word 3, $(subst -, ,$*))
|
||||
build/linux-mips-%: GOENV += GOMIPS=$(word 3, $(subst -, ,$*))
|
||||
@@ -182,8 +198,11 @@ build/linux-mips-softfloat/%: LDFLAGS += -s -w
|
||||
# boringcrypto
|
||||
build/linux-amd64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
||||
build/linux-arm64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
||||
build/linux-amd64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
||||
build/linux-arm64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
||||
|
||||
# fips140
|
||||
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
|
||||
GOOS=$(firstword $(subst -, , $*)) \
|
||||
@@ -214,10 +233,7 @@ vet:
|
||||
go vet $(VET_FLAGS) -v ./...
|
||||
|
||||
test:
|
||||
go test -v ./...
|
||||
|
||||
test-boringcrypto:
|
||||
GOEXPERIMENT=boringcrypto CGO_ENABLED=1 go test -ldflags "-checklinkname=0" -v ./...
|
||||
$(TEST_ENV) go test $(TEST_FLAGS) -v ./...
|
||||
|
||||
test-pkcs11:
|
||||
CGO_ENABLED=1 go test -v -tags pkcs11 ./...
|
||||
@@ -260,29 +276,75 @@ ifeq ($(words $(MAKECMDGOALS)),1)
|
||||
@$(MAKE) service ${.DEFAULT_GOAL} --no-print-directory
|
||||
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
|
||||
|
||||
smoke-docker: BUILD_ARGS += -race
|
||||
smoke-docker: GOENV += CGO_ENABLED=1
|
||||
smoke-docker: bin-docker
|
||||
cd .github/workflows/smoke/ && ./build.sh
|
||||
cd .github/workflows/smoke/ && ./smoke.sh
|
||||
cd .github/workflows/smoke/ && NAME="smoke-p256" CURVE="P256" ./build.sh
|
||||
cd .github/workflows/smoke/ && NAME="smoke-p256" ./smoke.sh
|
||||
# This is so we can limit `fips140` smoke test to just P256 curve.
|
||||
if [ "$(CURVE)" != "P256" ]; then cd .github/workflows/smoke/ && $(GOENV) ./build.sh; fi
|
||||
if [ "$(CURVE)" != "P256" ]; then cd .github/workflows/smoke/ && $(GOENV) ./smoke.sh; fi
|
||||
cd .github/workflows/smoke/ && $(GOENV) NAME="smoke-p256" CURVE="P256" ./build.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
|
||||
cd .github/workflows/smoke/ && ./build-relay.sh
|
||||
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
||||
cd .github/workflows/smoke/ && $(GOENV) ./build-relay.sh
|
||||
cd .github/workflows/smoke/ && $(GOENV) ./smoke-relay.sh
|
||||
|
||||
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
|
||||
smoke-docker-ipv6: smoke-docker
|
||||
|
||||
smoke-docker-race: BUILD_ARGS = -race
|
||||
smoke-docker-race: CGO_ENABLED = 1
|
||||
smoke-docker-race: smoke-docker
|
||||
smoke-self: bin
|
||||
cd .github/workflows/smoke/ && ./smoke-self.sh
|
||||
|
||||
smoke-vagrant/%: bin-docker build/%/nebula
|
||||
cd .github/workflows/smoke/ && ./build.sh $*
|
||||
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
||||
|
||||
.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 debug build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
||||
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin 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/%
|
||||
.DEFAULT_GOAL := bin
|
||||
|
||||
@@ -145,17 +145,27 @@ To build nebula for a specific platform (ex, Windows):
|
||||
|
||||
See the [Makefile](Makefile) for more details on build targets
|
||||
|
||||
## Curve P256 and BoringCrypto
|
||||
## Curve P256 and FIPS 140-3 mode
|
||||
|
||||
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.
|
||||
|
||||
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:
|
||||
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).
|
||||
|
||||
```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
|
||||
make bin-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.
|
||||
|
||||
## Credits
|
||||
|
||||
+30
-3
@@ -3,11 +3,14 @@ package main
|
||||
import (
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/fips140"
|
||||
"crypto/rand"
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"io"
|
||||
"math"
|
||||
"math/bits"
|
||||
"net/netip"
|
||||
"os"
|
||||
"strings"
|
||||
@@ -43,7 +46,28 @@ type caFlags struct {
|
||||
subnets *string
|
||||
}
|
||||
|
||||
func defaultCurve() string {
|
||||
if fips140.Enforced() {
|
||||
return "P256"
|
||||
}
|
||||
return "25519"
|
||||
}
|
||||
|
||||
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.set.Usage = func() {}
|
||||
cf.name = cf.set.String("name", "", "Required: name of the certificate authority")
|
||||
@@ -55,11 +79,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.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.argonMemory = cf.set.Uint("argon-memory", 2*1024*1024, "Optional: Argon2 memory parameter (in KiB) used for encrypted private key passphrase")
|
||||
cf.argonMemory = cf.set.Uint("argon-memory", defaultArgonMemory, "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.argonIterations = cf.set.Uint("argon-iterations", 1, "Optional: Argon2 iterations 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.encryption = cf.set.Bool("encrypt", false, "Optional: prompt for passphrase and write out-key in an encrypted format")
|
||||
cf.curve = cf.set.String("curve", "25519", "EdDSA/ECDSA Curve (25519, P256)")
|
||||
cf.curve = cf.set.String("curve", defaultCurve(), "EdDSA/ECDSA Curve (25519, P256)")
|
||||
cf.p11url = p11Flag(cf.set)
|
||||
|
||||
cf.ips = cf.set.String("ips", "", "Deprecated, see -networks")
|
||||
@@ -244,6 +268,9 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
||||
} else {
|
||||
switch *cf.curve {
|
||||
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
|
||||
pub, rawPriv, err = ed25519.GenerateKey(rand.Reader)
|
||||
if err != nil {
|
||||
|
||||
@@ -7,7 +7,9 @@ import (
|
||||
"bytes"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"math/bits"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -22,6 +24,18 @@ func Test_caSummary(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{}
|
||||
caHelp(ob)
|
||||
assert.Equal(
|
||||
@@ -29,9 +43,9 @@ func Test_caHelp(t *testing.T) {
|
||||
"Usage of "+os.Args[0]+" ca <flags>: create a self signed certificate authority\n"+
|
||||
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
|
||||
" -argon-iterations uint\n"+
|
||||
" \tOptional: Argon2 iterations parameter used for encrypted private key passphrase (default 1)\n"+
|
||||
" \tOptional: Argon2 iterations parameter used for encrypted private key passphrase (default "+defaultArgonIterations+")\n"+
|
||||
" -argon-memory uint\n"+
|
||||
" \tOptional: Argon2 memory parameter (in KiB) used for encrypted private key passphrase (default 2097152)\n"+
|
||||
" \tOptional: Argon2 memory parameter (in KiB) used for encrypted private key passphrase (default "+defaultArgonMemory+")\n"+
|
||||
" -argon-parallelism uint\n"+
|
||||
" \tOptional: Argon2 parallelism parameter used for encrypted private key passphrase (default 4)\n"+
|
||||
" -curve string\n"+
|
||||
@@ -188,10 +202,16 @@ func Test_ca(t *testing.T) {
|
||||
k, _ := pem.Decode(rb)
|
||||
ned, err := cert.UnmarshalNebulaEncryptedData(k.Bytes)
|
||||
require.NoError(t, err)
|
||||
// we won't know salt in advance, so just check start of string
|
||||
assert.Equal(t, uint32(2*1024*1024), ned.EncryptionMetadata.Argon2Parameters.Memory)
|
||||
|
||||
if bits.UintSize == 32 {
|
||||
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, uint32(1), ned.EncryptionMetadata.Argon2Parameters.Iterations)
|
||||
|
||||
// verify the key is valid and decrypt-able
|
||||
var curve cert.Curve
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
//go:build fips140enforce
|
||||
|
||||
//go:debug fips140=only
|
||||
|
||||
package main
|
||||
@@ -1,6 +1,8 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"crypto/fips140"
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"io"
|
||||
@@ -24,7 +26,7 @@ func newKeygenFlags() *keygenFlags {
|
||||
cf.set.Usage = func() {}
|
||||
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.curve = cf.set.String("curve", "25519", "ECDH Curve (25519, P256)")
|
||||
cf.curve = cf.set.String("curve", defaultCurve(), "ECDH Curve (25519, P256)")
|
||||
cf.p11url = p11Flag(cf.set)
|
||||
return &cf
|
||||
}
|
||||
@@ -61,6 +63,9 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
|
||||
} else {
|
||||
switch *cf.curve {
|
||||
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()
|
||||
curve = cert.Curve_CURVE25519
|
||||
case "P256":
|
||||
|
||||
@@ -2,6 +2,7 @@ package main
|
||||
|
||||
import (
|
||||
"crypto/ecdh"
|
||||
"crypto/fips140"
|
||||
"crypto/rand"
|
||||
"errors"
|
||||
"flag"
|
||||
@@ -268,6 +269,10 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
||||
}(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 != "" {
|
||||
var pubCurve cert.Curve
|
||||
rawPub, err := readInput("in-pub", *sf.inPubPath, &claims)
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
//go:build fips140enforce
|
||||
|
||||
//go:debug fips140=only
|
||||
|
||||
package main
|
||||
@@ -0,0 +1,96 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
cert_test "github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestControlStopClosesOnTimer reproduces the dnclient lifecycle: nebula runs as
|
||||
// a library, and on a config update dnclient calls Stop() in-process to tear the
|
||||
// old instance down before starting a new one. This boots a real nebula (real
|
||||
// blocking UDP sockets, tun disabled), lets it run, then Stop()s it on a timer
|
||||
// and asserts it actually closes. If the reader goroutines parked in recvmmsg
|
||||
// don't wake on Close(), Wait() blocks forever and this fails with a goroutine
|
||||
// dump instead of relying on a process signal to unstick them.
|
||||
func TestControlStopClosesOnTimer(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
dir := t.TempDir()
|
||||
|
||||
before := time.Now().Add(-time.Hour)
|
||||
after := time.Now().Add(time.Hour)
|
||||
ca, _, caKey, caPEM := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, before, after, nil, nil, nil)
|
||||
networks := []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}
|
||||
_, _, keyPEM, certPEM := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, "close-on-timer", before, after, networks, nil, nil)
|
||||
|
||||
caPath := filepath.Join(dir, "ca.pem")
|
||||
certPath := filepath.Join(dir, "cert.pem")
|
||||
keyPath := filepath.Join(dir, "key.pem")
|
||||
require.NoError(t, os.WriteFile(caPath, caPEM, 0o600))
|
||||
require.NoError(t, os.WriteFile(certPath, certPEM, 0o600))
|
||||
require.NoError(t, os.WriteFile(keyPath, keyPEM, 0o600))
|
||||
|
||||
// tun disabled so no device/root is needed; routines: 2 so we exercise the
|
||||
// multi-socket (SO_REUSEPORT) teardown, which is where dnclient runs.
|
||||
configBody := fmt.Sprintf(`
|
||||
pki:
|
||||
ca: %s
|
||||
cert: %s
|
||||
key: %s
|
||||
listen:
|
||||
host: 127.0.0.1
|
||||
port: 0
|
||||
tun:
|
||||
disabled: true
|
||||
firewall:
|
||||
outbound:
|
||||
- port: any
|
||||
proto: any
|
||||
host: any
|
||||
inbound:
|
||||
- port: any
|
||||
proto: any
|
||||
host: any
|
||||
routines: 2
|
||||
`, caPath, certPath, keyPath)
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dir, "config.yml"), []byte(configBody), 0o600))
|
||||
|
||||
c := config.NewC(l)
|
||||
require.NoError(t, c.Load(dir))
|
||||
|
||||
ctrl, err := nebula.Main(c, false, "close-on-timer", l, nil)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, ctrl.Start())
|
||||
|
||||
// Run like a live nebula, then close on a timer, exactly as dnclient does.
|
||||
<-time.NewTimer(5 * time.Second).C
|
||||
|
||||
stopped := make(chan struct{})
|
||||
go func() {
|
||||
ctrl.Stop() // closes the udp sockets (shutdown(2)) and the tun
|
||||
ctrl.Wait() // blocks until every reader goroutine has returned
|
||||
close(stopped)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-stopped:
|
||||
t.Log("nebula closed cleanly on timer")
|
||||
case <-time.After(10 * time.Second):
|
||||
buf := make([]byte, 1<<20)
|
||||
n := runtime.Stack(buf, true)
|
||||
t.Fatalf("nebula did NOT close within 10s of Stop(): a blocking reader never woke\n%s", buf[:n])
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
//go:build fips140enforce
|
||||
|
||||
//go:debug fips140=only
|
||||
|
||||
package main
|
||||
+34
-8
@@ -105,11 +105,18 @@ func (cm *connectionManager) getInactivityTimeout() time.Duration {
|
||||
}
|
||||
|
||||
func (cm *connectionManager) In(h *HostInfo) {
|
||||
h.in.Store(true)
|
||||
h.markIn()
|
||||
}
|
||||
|
||||
func (cm *connectionManager) Out(h *HostInfo) {
|
||||
h.out.Store(true)
|
||||
// OutNoRebind records outbound traffic without consuming the rebind epoch, for relayed sends: the direct path
|
||||
// to the relay consumes the edge, the via send must not.
|
||||
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) {
|
||||
@@ -128,8 +135,7 @@ func (cm *connectionManager) RelayUsed(localIndex uint32) {
|
||||
// getAndResetTrafficCheck returns if there was any inbound or outbound traffic within the last tick and
|
||||
// resets the state for this local index
|
||||
func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time) (bool, bool) {
|
||||
in := h.in.Swap(false)
|
||||
out := h.out.Swap(false)
|
||||
in, out := h.takeTraffic()
|
||||
if in || out {
|
||||
h.lastUsed = now
|
||||
}
|
||||
@@ -323,6 +329,12 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
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]]
|
||||
mainHostInfo := true
|
||||
if primary != nil && primary != hostinfo {
|
||||
@@ -340,7 +352,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
"tunnelCheck", m{"state": "alive", "method": "passive"},
|
||||
)
|
||||
}
|
||||
hostinfo.pendingDeletion.Store(false)
|
||||
hostinfo.setPendingDeletion(false)
|
||||
|
||||
if mainHostInfo {
|
||||
decision = tryRehandshake
|
||||
@@ -363,7 +375,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
return decision, hostinfo, primary
|
||||
}
|
||||
|
||||
if hostinfo.pendingDeletion.Load() {
|
||||
if hostinfo.isPendingDeletion() {
|
||||
// We have already sent a test packet and nothing was returned, this hostinfo is dead
|
||||
hostinfo.logger(cm.l).Info("Tunnel status",
|
||||
"tunnelCheck", m{"state": "dead", "method": "active"},
|
||||
@@ -414,7 +426,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
}
|
||||
}
|
||||
|
||||
hostinfo.pendingDeletion.Store(true)
|
||||
hostinfo.setPendingDeletion(true)
|
||||
cm.trafficTimer.Add(hostinfo.localIndexId, cm.pendingDeletionInterval)
|
||||
return decision, hostinfo, nil
|
||||
}
|
||||
@@ -448,6 +460,11 @@ func (cm *connectionManager) shouldSwapPrimary(current *HostInfo) bool {
|
||||
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())
|
||||
if crt == nil {
|
||||
//my cert was reloaded away. We should definitely swap from this tunnel
|
||||
@@ -544,6 +561,15 @@ func (cm *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
||||
"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)
|
||||
return
|
||||
}
|
||||
|
||||
+110
-36
@@ -25,6 +25,7 @@ func newTestLighthouse() *LightHouse {
|
||||
lighthouses := []netip.Addr{}
|
||||
staticList := map[netip.Addr]struct{}{}
|
||||
|
||||
lh.localAddrsFn = func(*LocalAllowList) []netip.Addr { return nil }
|
||||
lh.lighthouses.Store(&lighthouses)
|
||||
lh.staticList.Store(&staticList)
|
||||
|
||||
@@ -85,25 +86,25 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
||||
// We saw traffic out to vpnIp
|
||||
nc.Out(hostinfo)
|
||||
nc.In(hostinfo)
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
assert.True(t, hostinfo.out.Load())
|
||||
assert.True(t, hostinfo.in.Load())
|
||||
assert.True(t, hostinfo.sentSinceCheck())
|
||||
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
|
||||
// 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())
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
|
||||
// Do another traffic check tick, this host should be pending deletion now
|
||||
nc.Out(hostinfo)
|
||||
assert.True(t, hostinfo.out.Load())
|
||||
assert.True(t, hostinfo.sentSinceCheck())
|
||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||
assert.True(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.True(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
|
||||
@@ -167,37 +168,110 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
||||
// We saw traffic out to vpnIp
|
||||
nc.Out(hostinfo)
|
||||
nc.In(hostinfo)
|
||||
assert.True(t, hostinfo.in.Load())
|
||||
assert.True(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.True(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
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
|
||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
|
||||
// Do another traffic check tick, this host should be pending deletion now
|
||||
nc.Out(hostinfo)
|
||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||
assert.True(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.True(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
|
||||
// We saw traffic, should no longer be pending deletion
|
||||
nc.In(hostinfo)
|
||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
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) {
|
||||
l := test.NewLogger()
|
||||
localrange := netip.MustParsePrefix("10.1.1.1/24")
|
||||
@@ -252,31 +326,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
|
||||
nc.Out(hostinfo)
|
||||
nc.In(hostinfo)
|
||||
assert.True(t, hostinfo.out.Load())
|
||||
assert.True(t, hostinfo.in.Load())
|
||||
assert.True(t, hostinfo.sentSinceCheck())
|
||||
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
|
||||
now := time.Now()
|
||||
decision, _, _ := nc.makeTrafficDecision(hostinfo.localIndexId, now)
|
||||
assert.Equal(t, tryRehandshake, decision)
|
||||
assert.Equal(t, now, hostinfo.lastUsed)
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
|
||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*5))
|
||||
assert.Equal(t, doNothing, decision)
|
||||
assert.Equal(t, now, hostinfo.lastUsed)
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
|
||||
// Do another traffic check tick, should still not be pending deletion
|
||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*10))
|
||||
assert.Equal(t, doNothing, decision)
|
||||
assert.Equal(t, now, hostinfo.lastUsed)
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
|
||||
@@ -284,9 +358,9 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Minute*10))
|
||||
assert.Equal(t, closeTunnel, decision)
|
||||
assert.Equal(t, now, hostinfo.lastUsed)
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
}
|
||||
|
||||
+95
-3
@@ -2,15 +2,37 @@ package nebula
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/handshake"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
)
|
||||
|
||||
const ReplayWindow = 8192
|
||||
const (
|
||||
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 {
|
||||
eKey noiseutil.CipherState
|
||||
@@ -20,14 +42,22 @@ type ConnectionState struct {
|
||||
initiator bool
|
||||
messageCounter atomic.Uint64
|
||||
window *Bits
|
||||
decryptLock 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
|
||||
// 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
|
||||
// as missed traffic in the data plane.
|
||||
func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
|
||||
func newConnectionStateFromResult(r *handshake.Result) (*ConnectionState, error) {
|
||||
// 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{
|
||||
myCert: r.MyCert,
|
||||
initiator: r.Initiator,
|
||||
@@ -35,12 +65,13 @@ func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
|
||||
eKey: noiseutil.NewCipherState(r.EKey, r.Cipher),
|
||||
dKey: noiseutil.NewCipherState(r.DKey, r.Cipher),
|
||||
window: NewBits(ReplayWindow),
|
||||
epoch: sessionEpoch.Add(1),
|
||||
}
|
||||
ci.messageCounter.Add(r.MessageIndex)
|
||||
for i := uint64(1); i <= r.MessageIndex; i++ {
|
||||
ci.window.Update(nil, i)
|
||||
}
|
||||
return ci
|
||||
return ci, nil
|
||||
}
|
||||
|
||||
func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
||||
@@ -51,6 +82,67 @@ 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 {
|
||||
return cs.myCert.Curve()
|
||||
}
|
||||
|
||||
func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) ([]byte, error) {
|
||||
cs.decryptLock.Lock()
|
||||
result := cs.window.Check(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return nil, ErrAlreadySeen
|
||||
}
|
||||
|
||||
out, err := cs.dKey.DecryptDanger(packet[header.Len:header.Len], packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cs.decryptLock.Lock()
|
||||
result = cs.window.Update(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return nil, ErrAlreadySeen
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) error {
|
||||
cs.decryptLock.Lock()
|
||||
result := cs.window.Check(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return ErrAlreadySeen
|
||||
}
|
||||
|
||||
// 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()]
|
||||
signatureValue := packet[len(packet)-cs.dKey.Overhead():]
|
||||
_, err := cs.dKey.DecryptDanger(nil, signedPayload, signatureValue, messageCounter, nb)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cs.decryptLock.Lock()
|
||||
result = cs.window.Update(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return ErrAlreadySeen
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -6,10 +6,13 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
ct "github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/handshake"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
@@ -79,11 +82,77 @@ func runTestHandshake(t *testing.T) (initR, respR *handshake.Result) {
|
||||
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) {
|
||||
initR, respR := runTestHandshake(t)
|
||||
|
||||
t.Run("initiator", func(t *testing.T) {
|
||||
ci := newConnectionStateFromResult(initR)
|
||||
ci, err := newConnectionStateFromResult(initR)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, ci.initiator)
|
||||
assert.Equal(t, initR.MyCert, ci.myCert)
|
||||
assert.Equal(t, initR.RemoteCert, ci.peerCert)
|
||||
@@ -102,8 +171,17 @@ func TestNewConnectionStateFromResult(t *testing.T) {
|
||||
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) {
|
||||
ci := newConnectionStateFromResult(respR)
|
||||
ci, err := newConnectionStateFromResult(respR)
|
||||
require.NoError(t, err)
|
||||
assert.False(t, ci.initiator)
|
||||
assert.Equal(t, respR.MyCert, ci.myCert)
|
||||
assert.Equal(t, respR.RemoteCert, ci.peerCert)
|
||||
|
||||
+11
-3
@@ -53,6 +53,7 @@ type Control struct {
|
||||
statsStart func()
|
||||
dnsStart func()
|
||||
lighthouseStart func()
|
||||
networkChangeStart func(rebind func())
|
||||
connectionManagerStart func(context.Context)
|
||||
}
|
||||
|
||||
@@ -104,6 +105,9 @@ func (c *Control) Start() error {
|
||||
if c.dnsStart != nil {
|
||||
go c.dnsStart()
|
||||
}
|
||||
if c.networkChangeStart != nil {
|
||||
go c.networkChangeStart(c.RebindUDPServer)
|
||||
}
|
||||
if c.connectionManagerStart != nil {
|
||||
go c.connectionManagerStart(c.ctx)
|
||||
}
|
||||
@@ -111,7 +115,7 @@ func (c *Control) Start() error {
|
||||
c.lighthouseStart()
|
||||
}
|
||||
|
||||
c.f.triggerShutdown = c.Stop
|
||||
c.f.triggerShutdown = func() { go c.Stop() }
|
||||
|
||||
// Start reading packets.
|
||||
c.f.run()
|
||||
@@ -198,13 +202,17 @@ func (c *Control) RebindUDPServer() {
|
||||
return
|
||||
}
|
||||
|
||||
_ = c.f.outside.Rebind()
|
||||
// A failure here means we are likely still pinned to the interface we came up on, so the rest of this is
|
||||
// unlikely to help. Say so instead of silently carrying on as if we rebound.
|
||||
if err := c.f.outside.Rebind(); err != nil {
|
||||
c.l.Error("Failed to rebind udp socket", "error", err)
|
||||
}
|
||||
|
||||
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
|
||||
c.f.lightHouse.SendUpdate()
|
||||
|
||||
// Let the main interface know that we rebound so that underlying tunnels know to trigger punches from their remotes
|
||||
c.f.rebindCount++
|
||||
c.f.rebindEpoch.Add(1)
|
||||
}
|
||||
|
||||
// ListHostmapHosts returns details about the actual or pending (handshaking) hostmap by vpn ip
|
||||
|
||||
@@ -78,7 +78,7 @@ func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
|
||||
inside: dev,
|
||||
outside: conn,
|
||||
writers: []udp.Conn{conn},
|
||||
batchers: make([]batch.RxBatcher, 1),
|
||||
batchers: make([]*batch.MultiCoalescer, 1),
|
||||
routines: 1,
|
||||
hostMap: newHostMap(l),
|
||||
lightHouse: lh,
|
||||
@@ -148,8 +148,8 @@ func (c *fakeConn) Rebind() error { c.rebinds++; ret
|
||||
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||
func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil }
|
||||
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
||||
func (c *fakeConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
|
||||
return nil
|
||||
func (c *fakeConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) {
|
||||
return len(bufs), nil
|
||||
}
|
||||
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
||||
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
||||
@@ -177,7 +177,7 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
||||
inside: dev,
|
||||
outside: conn,
|
||||
writers: []udp.Conn{conn},
|
||||
batchers: make([]batch.RxBatcher, 2),
|
||||
batchers: make([]*batch.MultiCoalescer, 2),
|
||||
routines: 2,
|
||||
l: test.NewLogger(),
|
||||
}
|
||||
|
||||
+23
-1
@@ -108,7 +108,29 @@ func (c *Control) GetVpnAddrs() []netip.Addr {
|
||||
}
|
||||
|
||||
func (c *Control) GetUDPAddr() netip.AddrPort {
|
||||
return c.f.outside.(*udp.TesterConn).Addr
|
||||
return c.f.outside.(*udp.TesterConn).GetAddr()
|
||||
}
|
||||
|
||||
// SetUDPAddr moves this node to a new underlay address, standing in for a laptop waking up on a different
|
||||
// network. Register the new address with the router as well or nothing will route back.
|
||||
func (c *Control) SetUDPAddr(addr netip.AddrPort) {
|
||||
c.f.outside.(*udp.TesterConn).SetAddr(addr)
|
||||
}
|
||||
|
||||
// SetLocalAddrsFn replaces underlay address discovery so a test can advertise its simulated address instead of
|
||||
// whatever this machine's NICs happen to be. Call it before Start, SendUpdate reads it from the update worker.
|
||||
func (c *Control) SetLocalAddrsFn(fn func(*LocalAllowList) []netip.Addr) {
|
||||
c.f.lightHouse.localAddrsFn = fn
|
||||
}
|
||||
|
||||
// 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 {
|
||||
|
||||
@@ -0,0 +1,187 @@
|
||||
// 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)
|
||||
}
|
||||
@@ -0,0 +1,171 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,154 @@
|
||||
//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)))
|
||||
}
|
||||
@@ -0,0 +1,163 @@
|
||||
//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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
//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, ""
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
//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
|
||||
}
|
||||
@@ -0,0 +1,111 @@
|
||||
//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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
//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)
|
||||
}
|
||||
+16
-7
@@ -97,8 +97,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
||||
newAddr := getDnsServerAddr(c)
|
||||
|
||||
d.serverMu.Lock()
|
||||
running := d.server
|
||||
runningStarted := d.started
|
||||
running := d.server != nil
|
||||
sameAddr := d.addr == newAddr
|
||||
d.addr = newAddr
|
||||
d.enabled.Store(enabled)
|
||||
@@ -112,7 +111,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
||||
}
|
||||
|
||||
if !enabled {
|
||||
if running != nil {
|
||||
if running {
|
||||
d.Stop()
|
||||
}
|
||||
// Drop any records that accumulated while enabled; a later re-enable
|
||||
@@ -121,12 +120,12 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
if running == nil {
|
||||
if !running {
|
||||
// Was disabled (or never started); bring it up now.
|
||||
go d.Start()
|
||||
} else if !sameAddr {
|
||||
d.shutdownServer(running, runningStarted, "reload")
|
||||
// Old Start goroutine has now exited; bring up a fresh listener on the new address.
|
||||
// Stop clears the slot before shutting down, otherwise the Start below can find the dying server and refuse
|
||||
d.Stop()
|
||||
go d.Start()
|
||||
}
|
||||
|
||||
@@ -162,7 +161,9 @@ func (d *dnsServer) Start() {
|
||||
|
||||
started := make(chan struct{})
|
||||
d.serverMu.Lock()
|
||||
if d.ctx.Err() != nil {
|
||||
// Re-check enabled under the lock, a disable that raced our check above snapshots the slot under it too.
|
||||
// Two reloads in quick succession can both spawn a Start, the loser would orphan the live listener past Stop
|
||||
if d.ctx.Err() != nil || d.server != nil || !d.enabled.Load() {
|
||||
d.serverMu.Unlock()
|
||||
return
|
||||
}
|
||||
@@ -200,6 +201,14 @@ func (d *dnsServer) Start() {
|
||||
close(started)
|
||||
}
|
||||
|
||||
// Release our slot, unless a reload already replaced us, so a dead listener can't block a future Start
|
||||
d.serverMu.Lock()
|
||||
if d.server == server {
|
||||
d.server = nil
|
||||
d.started = nil
|
||||
}
|
||||
d.serverMu.Unlock()
|
||||
|
||||
if err != nil {
|
||||
d.l.Warn("Failed to run the DNS responder", "error", err)
|
||||
}
|
||||
|
||||
+206
-4
@@ -194,14 +194,51 @@ func TestDnsServer_reload_initial_serveDnsWithoutLighthouse(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDnsServer_reload_sameAddr_noOp(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", "0", true, true)
|
||||
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
// No server running yet, no addr change. Reload should not spawn anything.
|
||||
|
||||
go ds.Start()
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
before := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
require.NotNil(t, before)
|
||||
|
||||
// Same address, so the running listener must be left alone rather than rebuilt under live queries
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
assert.True(t, ds.enabled.Load())
|
||||
assert.Nil(t, ds.server)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
after := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
assert.Same(t, before, after, "a same-address reload must not restart the listener")
|
||||
|
||||
ds.Stop()
|
||||
}
|
||||
|
||||
// The branch the old sameAddr test was accidentally hitting: enabled with nothing running means reload starts it.
|
||||
func TestDnsServer_reload_whenNotRunning_starts(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
|
||||
// initial only records config, it never starts anything
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
ds.serverMu.Lock()
|
||||
assert.Nil(t, ds.server, "the initial reload must not start a listener")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.NotNil(t, ds.server, "a reload with nothing running should bring DNS up")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
ds.Stop()
|
||||
}
|
||||
|
||||
func TestDnsServer_StartStop_lifecycle(t *testing.T) {
|
||||
@@ -427,3 +464,168 @@ func waitFor(t *testing.T, cond func() bool) {
|
||||
}
|
||||
t.Fatal("timed out waiting for condition")
|
||||
}
|
||||
|
||||
// Two reloads in quick succession, or a HUP before Control.Start, can race two Starts at the same listener.
|
||||
func TestDnsServer_Start_isIdempotent(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
go ds.Start()
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
first := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
require.NotNil(t, first)
|
||||
|
||||
// If the second Start replaces the tracked server, Stop kills the wrong one and the port leaks
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
ds.Start()
|
||||
close(done)
|
||||
}()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("second Start never returned")
|
||||
}
|
||||
|
||||
ds.serverMu.Lock()
|
||||
second := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
assert.Same(t, first, second, "a second Start must not replace the running server")
|
||||
|
||||
// The real proof, after Stop the port must actually be free
|
||||
ds.Stop()
|
||||
waitFor(t, func() bool {
|
||||
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
_ = pc.Close()
|
||||
return true
|
||||
})
|
||||
}
|
||||
|
||||
// An address change must actually end up listening on the new port. Start's guard refuses when a server is already
|
||||
// installed, so reload has to clear the slot before shutting the old one down.
|
||||
func TestDnsServer_reload_addrChange_restarts(t *testing.T) {
|
||||
first := freeUDPPort(t)
|
||||
second := freeUDPPort(t)
|
||||
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", first, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
go ds.Start()
|
||||
waitForBind(t, ds)
|
||||
|
||||
// Cycle a few times, the failure this guards against depends on which goroutine wins serverMu
|
||||
for i := range 8 {
|
||||
want := second
|
||||
if i%2 == 1 {
|
||||
want = first
|
||||
}
|
||||
setDnsConfig(c, "127.0.0.1", want, true, true)
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
srv := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
require.NotNil(t, srv, "reload left DNS down instead of restarting it")
|
||||
require.Equal(t, "127.0.0.1:"+want, srv.Addr, "reload should be serving the new address")
|
||||
}
|
||||
|
||||
// Land back on second so the port assertions below are meaningful
|
||||
setDnsConfig(c, "127.0.0.1", second, true, true)
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
// The old port must be released and the new one actually held
|
||||
waitFor(t, func() bool {
|
||||
pc, err := net.ListenPacket("udp", "127.0.0.1:"+first)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
_ = pc.Close()
|
||||
return true
|
||||
})
|
||||
_, err := net.ListenPacket("udp", "127.0.0.1:"+second)
|
||||
require.Error(t, err, "the new address should be bound by the DNS responder")
|
||||
|
||||
ds.Stop()
|
||||
}
|
||||
|
||||
// A listener that dies on its own must release the slot, or a later same-addr reload sees it as running and no-ops.
|
||||
func TestDnsServer_Start_bindFailure_releasesSlot(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
blocker, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||
require.NoError(t, err)
|
||||
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
ds.Start() // returns once the bind fails
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.Nil(t, ds.server, "a listener that failed to bind must not stay parked in the slot")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
// With the slot released, a reload can retry once the port frees up
|
||||
require.NoError(t, blocker.Close())
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.NotNil(t, ds.server, "a same-addr reload should retry after a failed bind")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
ds.Stop()
|
||||
}
|
||||
|
||||
// A disable that lands while Start is between its unlocked check and the guard must not leave a listener behind.
|
||||
func TestDnsServer_Start_refusesWhenDisabledUnderLock(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
require.True(t, ds.enabled.Load())
|
||||
|
||||
// Holding serverMu parks Start on the lock, the only way to land the disable in that window on purpose
|
||||
ds.serverMu.Lock()
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
ds.Start()
|
||||
close(done)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
ds.serverMu.Unlock()
|
||||
t.Fatal("Start returned early, the test never exercised the window")
|
||||
case <-time.After(time.Millisecond * 100):
|
||||
}
|
||||
|
||||
// The disable reload's critical section. It sees nothing running, so it never calls Stop.
|
||||
ds.enabled.Store(false)
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("Start never returned")
|
||||
}
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.Nil(t, ds.server, "Start must not install a listener a disable already cancelled")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||
require.NoError(t, err, "an orphaned listener is still holding the port")
|
||||
_ = pc.Close()
|
||||
}
|
||||
|
||||
@@ -725,6 +725,70 @@ func TestReestablishRelays(t *testing.T) {
|
||||
|
||||
}
|
||||
|
||||
func TestRelayHandshakeOverDisestablishedEntry(t *testing.T) {
|
||||
t.Parallel()
|
||||
// If them tears down the tunnel while me keeps Established relay state, me's next
|
||||
// handshake flows through the relay with no fresh CreateRelayRequest and lands on
|
||||
// them's Disestablished terminal relay entry. them must re-establish that entry, or
|
||||
// its first transmit deletes its only relay and the tunnel is born transmit-dead:
|
||||
// them can receive but every send is silently dropped.
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them ", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
|
||||
|
||||
// Teach my how to get to the relay and that their can be reached via the relay
|
||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||
|
||||
// Build a router so we don't have to reason who gets which packet
|
||||
r := router.NewR(t, myControl, relayControl, theirControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
// Start the servers
|
||||
myControl.Start()
|
||||
relayControl.Start()
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger a handshake from me to them via the relay")
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||
oldIdx := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false).LocalIndex
|
||||
|
||||
t.Log("Close the tunnel on them only, marking their relay entry Disestablished")
|
||||
theirControl.CloseTunnel(myVpnIpNet[0].Addr(), true)
|
||||
|
||||
t.Log("Re-handshake from me, riding the still-Established relay state")
|
||||
myControl.ReHandshake(theirVpnIpNet[0].Addr())
|
||||
for {
|
||||
h := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
|
||||
if h != nil && h.LocalIndex != oldIdx && h.RemoteIndex != 0 {
|
||||
break
|
||||
}
|
||||
r.RouteForAllExitFunc(func(*udp.Packet, *nebula.Control) router.ExitType {
|
||||
return router.RouteAndExit
|
||||
})
|
||||
}
|
||||
|
||||
hAtThem := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
||||
require.NotNil(t, hAtThem, "them should have completed the relayed handshake")
|
||||
require.Equal(t, []netip.Addr{relayVpnIpNet[0].Addr()}, hAtThem.CurrentRelaysToMe, "them should know a relay for the new tunnel")
|
||||
|
||||
t.Log("Send from them to me; their only relay entry must survive the transmit")
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||
require.Never(t, func() bool {
|
||||
h := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
||||
return h == nil || len(h.CurrentRelaysToMe) == 0
|
||||
}, time.Second, 10*time.Millisecond, "them deleted its only relay entry; the tunnel is permanently transmit-dead")
|
||||
|
||||
p = r.RouteForAllUntilTxTun(myControl)
|
||||
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
||||
r.RenderHostmaps("Final hostmaps", myControl, relayControl, theirControl)
|
||||
}
|
||||
|
||||
func TestStage1RaceRelays(t *testing.T) {
|
||||
t.Parallel()
|
||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||
|
||||
@@ -0,0 +1,282 @@
|
||||
//go:build e2e_testing
|
||||
// +build e2e_testing
|
||||
|
||||
package e2e
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/e2e/router"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// reportedAddrs is what the lighthouse would hand a peer asking where vpnAddr is.
|
||||
func reportedAddrs(t *testing.T, lh *nebula.Control, vpnAddr netip.Addr) []netip.AddrPort {
|
||||
t.Helper()
|
||||
cm := lh.QueryLighthouse(vpnAddr)
|
||||
if cm == nil {
|
||||
return nil
|
||||
}
|
||||
var out []netip.AddrPort
|
||||
for _, c := range *cm {
|
||||
out = append(out, c.Reported...)
|
||||
out = append(out, c.Learned...)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// waitForLighthouseMsg routes until a lighthouse message lands on lh, or gives up. Reports whether one arrived.
|
||||
func waitForLighthouseMsg(t *testing.T, r *router.R, lh *nebula.Control, wait time.Duration) bool {
|
||||
t.Helper()
|
||||
h := &header.H{}
|
||||
return r.RouteForAllExitFuncOrTimeout(wait, func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||
if c != lh {
|
||||
return router.KeepRouting
|
||||
}
|
||||
// Punches are a single byte and never parse, they are just not what we are after
|
||||
if err := h.Parse(p.Data); err != nil {
|
||||
return router.KeepRouting
|
||||
}
|
||||
if h.Type == header.LightHouse {
|
||||
return router.RouteAndExit
|
||||
}
|
||||
return router.KeepRouting
|
||||
})
|
||||
}
|
||||
|
||||
// A laptop that changes networks has to tell the lighthouse promptly, otherwise the lighthouse keeps handing peers
|
||||
// the old address and their punches land nowhere. On a long lighthouse interval the only thing that closes that
|
||||
// window is the rebind, which on darwin the network change monitor drives. The e2e build compiles the monitor out,
|
||||
// so we call RebindUDPServer directly, which is the same thing the monitor does.
|
||||
func TestRebindSendsLighthouseUpdate(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||
"lighthouse": m{"am_lighthouse": true},
|
||||
})
|
||||
|
||||
// 600s interval, so nothing scheduled can send an update during this test. A rebind is the only thing that can.
|
||||
myControl, _, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
||||
"lighthouse": m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
},
|
||||
"static_host_map": m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
},
|
||||
})
|
||||
|
||||
r := router.NewR(t, lhControl, myControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
|
||||
// Let the startup registration finish, then clear everything it left behind
|
||||
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
|
||||
// Nothing should be talking to the lighthouse on its own now
|
||||
require.False(t, waitForLighthouseMsg(t, r, lhControl, time.Millisecond*200),
|
||||
"nothing should reach the lighthouse before the rebind")
|
||||
|
||||
myControl.RebindUDPServer()
|
||||
|
||||
assert.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5),
|
||||
"a rebind should push an update to the lighthouse rather than waiting out the interval")
|
||||
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
}
|
||||
|
||||
// The other half of a rebind: every live tunnel requeries the lighthouse on its next send. That query is what makes
|
||||
// the lighthouse tell the peer to punch toward our new address, which is the part that actually revives a tunnel
|
||||
// whose remote NAT state died while we were on a different network.
|
||||
func TestRebindRequeriesPeersOnNextSend(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||
"lighthouse": m{"am_lighthouse": true},
|
||||
})
|
||||
|
||||
lhCfg := m{
|
||||
"lighthouse": m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
// Without this the peers advertise this machine's real addresses and then try to punch at them,
|
||||
// which the router has no route for.
|
||||
"local_allow_list": m{
|
||||
"10.0.0.0/24": true,
|
||||
"::/0": false,
|
||||
},
|
||||
},
|
||||
"static_host_map": m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
},
|
||||
}
|
||||
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", lhCfg)
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", lhCfg)
|
||||
|
||||
r := router.NewR(t, lhControl, myControl, theirControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
theirControl.Start()
|
||||
r.RouteFor(time.Millisecond * 500)
|
||||
|
||||
// Point the peers at each other directly, this test is about the rebind and not about lighthouse discovery
|
||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("initial")))
|
||||
r.RouteFor(time.Second)
|
||||
require.NotNil(t, myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false), "expected a tunnel to them")
|
||||
r.RouteFor(time.Millisecond * 300)
|
||||
|
||||
// Assert on what the peer sees rather than on lighthouse traffic. A query for them makes the lighthouse send
|
||||
// them a punch notification, which is the whole point. Our own update to the lighthouse sends them nothing,
|
||||
// so this cannot be satisfied by the update the rebind itself pushes.
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("quiet")))
|
||||
require.False(t, waitForLighthouseMsg(t, r, theirControl, time.Millisecond*300),
|
||||
"an ordinary send should not requery the lighthouse")
|
||||
|
||||
myControl.RebindUDPServer()
|
||||
r.RouteFor(time.Millisecond * 300) // let the update the rebind itself sends pass by
|
||||
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("after rebind")))
|
||||
assert.True(t, waitForLighthouseMsg(t, r, theirControl, time.Second*5),
|
||||
"the first send after a rebind should requery the lighthouse, which then tells the peer to punch at us")
|
||||
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
theirControl.Stop()
|
||||
}
|
||||
|
||||
// The scenario this whole thing exists for: a laptop sleeps at the office and wakes up at home on a new address.
|
||||
// Until it tells the lighthouse, the lighthouse keeps handing peers the office address, so their punches land
|
||||
// nowhere and the tunnel stays dead. On a long interval the rebind is the only thing that closes that window.
|
||||
func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||
"lighthouse": m{"am_lighthouse": true},
|
||||
})
|
||||
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
||||
"lighthouse": m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
},
|
||||
"static_host_map": m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
},
|
||||
})
|
||||
|
||||
// Advertise wherever we currently are rather than this machine's real NICs, read fresh each time so a move
|
||||
// is picked up.
|
||||
myControl.SetLocalAddrsFn(func(*nebula.LocalAllowList) []netip.Addr {
|
||||
return []netip.Addr{myControl.GetUDPAddr().Addr()}
|
||||
})
|
||||
|
||||
r := router.NewR(t, lhControl, myControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
|
||||
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
|
||||
require.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), myUdpAddr,
|
||||
"the lighthouse should know the address we started on")
|
||||
|
||||
// Wake up somewhere else
|
||||
newAddr := netip.MustParseAddrPort("10.0.0.99:4242")
|
||||
myControl.SetUDPAddr(newAddr)
|
||||
r.AddRoute(newAddr.Addr(), newAddr.Port(), myControl)
|
||||
|
||||
// Nothing has told the lighthouse, and with interval 600 nothing scheduled will
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
require.NotContains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
|
||||
"the lighthouse should still be handing out the old address before the rebind")
|
||||
|
||||
myControl.RebindUDPServer()
|
||||
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an update after the rebind")
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
|
||||
assert.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
|
||||
"after the rebind the lighthouse should hand peers our new address")
|
||||
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
}
|
||||
|
||||
// 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()
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
//go:build e2e_testing
|
||||
// +build e2e_testing
|
||||
|
||||
package e2e
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/e2e/router"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
)
|
||||
|
||||
// TestRecoveryTiming measures how long a tunnel takes to come back after the peer stops accepting our traffic,
|
||||
// which is what a laptop waking on a new network looks like from the peer's side: its NAT has no state for where
|
||||
// we are now, so everything we send disappears.
|
||||
//
|
||||
// It is a measurement, not a pass/fail assertion. Recovery is timed to the moment the peer punches back at us,
|
||||
// since that is when its NAT opens and the tunnel is usable again.
|
||||
//
|
||||
// go test -tags e2e_testing -v -run TestRecoveryTiming ./e2e/
|
||||
func TestRecoveryTiming(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
rebind bool
|
||||
}{
|
||||
{"no trigger", false},
|
||||
{"rebind counter", true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
d, lost := measureRecovery(t, tc.rebind)
|
||||
t.Logf("RESULT %-16s recovered in %-9v (%d packets lost)", tc.name, d.Round(time.Millisecond), lost)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// measureRecovery returns how long until the peer punched back, and how many of our packets died meanwhile. When
|
||||
// rebind is true we call RebindUDPServer once the tunnel goes dark, which is what the darwin network change
|
||||
// monitor does and what iOS has always done. When false, nothing tells nebula anything is wrong.
|
||||
func measureRecovery(t *testing.T, rebind bool) (time.Duration, int) {
|
||||
t.Helper()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||
"lighthouse": m{"am_lighthouse": true},
|
||||
})
|
||||
|
||||
peerCfg := m{
|
||||
"lighthouse": m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
"local_allow_list": m{
|
||||
"10.0.0.0/24": true,
|
||||
"::/0": false,
|
||||
},
|
||||
},
|
||||
"static_host_map": m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
},
|
||||
}
|
||||
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", peerCfg)
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", peerCfg)
|
||||
|
||||
r := router.NewR(t, lhControl, myControl, theirControl)
|
||||
defer r.RenderFlow()
|
||||
defer func() {
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
theirControl.Stop()
|
||||
}()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
theirControl.Start()
|
||||
r.RouteFor(time.Millisecond * 500)
|
||||
|
||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("establish")))
|
||||
r.RouteFor(time.Second)
|
||||
if myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false) == nil {
|
||||
t.Fatal("failed to establish the tunnel we are measuring")
|
||||
}
|
||||
r.RouteFor(time.Millisecond * 500)
|
||||
|
||||
// From here the peer's NAT has no state for us, everything we send it disappears
|
||||
start := time.Now()
|
||||
blackholed := 0
|
||||
var recovered time.Duration
|
||||
|
||||
if rebind {
|
||||
myControl.RebindUDPServer()
|
||||
}
|
||||
|
||||
// Keep the tun busy the way someone retrying a stalled connection would
|
||||
stop := make(chan struct{})
|
||||
defer close(stop)
|
||||
go func() {
|
||||
tick := time.NewTicker(time.Millisecond * 200)
|
||||
defer tick.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
case <-tick.C:
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(
|
||||
theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("retry")))
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
r.RouteForAllExitFuncOrTimeout(time.Second*30, func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||
if c == theirControl && p.From == myControl.GetUDPAddr() {
|
||||
blackholed++
|
||||
return router.Drop
|
||||
}
|
||||
|
||||
// The peer reaching us directly is the moment its NAT opened, whether that is a punch or a handshake
|
||||
if c == myControl && p.From == theirUdpAddr {
|
||||
recovered = time.Since(start)
|
||||
return router.RouteAndExit
|
||||
}
|
||||
|
||||
return router.KeepRouting
|
||||
})
|
||||
|
||||
if recovered == 0 {
|
||||
t.Fatalf("no recovery within 30s (%d packets blackholed)", blackholed)
|
||||
}
|
||||
return recovered, blackholed
|
||||
}
|
||||
+151
-31
@@ -6,11 +6,13 @@ package router
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"maps"
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"slices"
|
||||
"sort"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
@@ -22,7 +24,6 @@ import (
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"golang.org/x/exp/maps"
|
||||
)
|
||||
|
||||
// outNatKey is the (from, to) pair used by outNat. Comparable struct, so it works as a map key without the
|
||||
@@ -114,6 +115,28 @@ type packet struct {
|
||||
packet *udp.Packet
|
||||
tun bool // a packet pulled off a tun device
|
||||
rx bool // the packet was received by a udp device
|
||||
|
||||
// h is the nebula header, parsed once when the packet is recorded. parseErr says why there isn't one, which
|
||||
// the flow log reports rather than hiding. Punchy sends a single byte, so an unparseable packet is normal.
|
||||
h header.H
|
||||
parseErr error
|
||||
}
|
||||
|
||||
// fromAddr and toAddr are the addresses this packet actually travelled between. Reading them off the control
|
||||
// instead would misreport the whole history once a test moves a node. Tun packets are synthesized without
|
||||
// addresses, so they fall back to the control.
|
||||
func (p *packet) fromAddr() netip.AddrPort {
|
||||
if p.tun || !p.packet.From.IsValid() {
|
||||
return p.from.GetUDPAddr()
|
||||
}
|
||||
return p.packet.From
|
||||
}
|
||||
|
||||
func (p *packet) toAddr() netip.AddrPort {
|
||||
if p.tun || !p.packet.To.IsValid() {
|
||||
return p.to.GetUDPAddr()
|
||||
}
|
||||
return p.packet.To
|
||||
}
|
||||
|
||||
func (p *packet) WasReceived() {
|
||||
@@ -131,6 +154,9 @@ const (
|
||||
ExitNow ExitType = 1
|
||||
// RouteAndExit routes this packet and exits immediately afterwards
|
||||
RouteAndExit ExitType = 2
|
||||
// Drop discards this packet without delivering it and keeps routing. Use it to simulate a blackhole, such as
|
||||
// a restrictive NAT refusing traffic from an address it has not seen.
|
||||
Drop ExitType = 3
|
||||
)
|
||||
|
||||
type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
|
||||
@@ -141,7 +167,9 @@ type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
|
||||
func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
if err := os.MkdirAll("mermaid", 0755); err != nil {
|
||||
// t.Name() contains a slash for subtests, so the flow log can land in a nested directory
|
||||
fn := filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name()))
|
||||
if err := os.MkdirAll(filepath.Dir(fn), 0755); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
@@ -152,7 +180,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
||||
outNat: make(map[outNatKey]netip.AddrPort),
|
||||
flow: []flowEntry{},
|
||||
ignoreFlows: []ignoreFlow{},
|
||||
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
||||
fn: fn,
|
||||
t: t,
|
||||
cancelRender: cancel,
|
||||
}
|
||||
@@ -249,7 +277,7 @@ func (r *R) renderFlow() {
|
||||
continue
|
||||
}
|
||||
|
||||
addr := e.packet.from.GetUDPAddr()
|
||||
addr := e.packet.fromAddr()
|
||||
if _, ok := participants[addr]; ok {
|
||||
continue
|
||||
}
|
||||
@@ -268,7 +296,6 @@ func (r *R) renderFlow() {
|
||||
}
|
||||
|
||||
// Print packets
|
||||
h := &header.H{}
|
||||
for _, e := range r.flow {
|
||||
if e.packet == nil {
|
||||
//fmt.Fprintf(f, " note over %s: %s\n", strings.Join(participantsVals, ", "), e.note)
|
||||
@@ -280,21 +307,22 @@ func (r *R) renderFlow() {
|
||||
fmt.Fprintln(f, r.formatUdpPacket(p))
|
||||
|
||||
} else {
|
||||
if err := h.Parse(p.packet.Data); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
line := "--x"
|
||||
if p.rx {
|
||||
line = "->>"
|
||||
}
|
||||
|
||||
fmt.Fprintf(f,
|
||||
" %s%s%s: %s(%s), index %v, counter: %v\n",
|
||||
normalizeName(p.from.GetUDPAddr().String()),
|
||||
detail := fmt.Sprintf("%s(%s), index %v, counter: %v",
|
||||
p.h.TypeName(), p.h.SubTypeName(), p.h.RemoteIndex, p.h.MessageCounter)
|
||||
if p.parseErr != nil {
|
||||
detail = fmt.Sprintf("unparsed, %v (%d bytes)", p.parseErr, len(p.packet.Data))
|
||||
}
|
||||
|
||||
fmt.Fprintf(f, " %s%s%s: %s\n",
|
||||
normalizeName(p.fromAddr().String()),
|
||||
line,
|
||||
normalizeName(p.to.GetUDPAddr().String()),
|
||||
h.TypeName(), h.SubTypeName(), h.RemoteIndex, h.MessageCounter,
|
||||
normalizeName(p.toAddr().String()),
|
||||
detail,
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -347,7 +375,7 @@ func (r *R) RenderHostmaps(title string, controls ...*nebula.Control) {
|
||||
}
|
||||
|
||||
func (r *R) renderHostmaps(title string) {
|
||||
c := maps.Values(r.controls)
|
||||
c := slices.AppendSeq(make([]*nebula.Control, 0, len(r.controls)), maps.Values(r.controls))
|
||||
sort.SliceStable(c, func(i, j int) bool {
|
||||
return c[i].GetVpnAddrs()[0].Compare(c[j].GetVpnAddrs()[0]) > 0
|
||||
})
|
||||
@@ -408,29 +436,34 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
|
||||
|
||||
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
|
||||
|
||||
if len(r.ignoreFlows) > 0 {
|
||||
var h header.H
|
||||
err := h.Parse(p.Data)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
var h header.H
|
||||
var parseErr error
|
||||
if !tun {
|
||||
parseErr = h.Parse(p.Data)
|
||||
}
|
||||
|
||||
for _, i := range r.ignoreFlows {
|
||||
if !tun {
|
||||
if i.messageType == h.Type && i.subType == h.Subtype {
|
||||
return nil
|
||||
}
|
||||
} else if i.tun.HasValue && i.tun.IsTrue {
|
||||
// Decide before copying, the copy comes from a freelist and an ignored packet would never be released
|
||||
for _, i := range r.ignoreFlows {
|
||||
if tun {
|
||||
if i.tun.HasValue && i.tun.IsTrue {
|
||||
return nil
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// A packet we could not parse has no type to match against, so no rule can ignore it
|
||||
if parseErr == nil && i.messageType == h.Type && i.subType == h.Subtype {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
fp := &packet{
|
||||
from: from,
|
||||
to: to,
|
||||
packet: p.Copy(),
|
||||
tun: tun,
|
||||
from: from,
|
||||
to: to,
|
||||
packet: p.Copy(),
|
||||
tun: tun,
|
||||
h: h,
|
||||
parseErr: parseErr,
|
||||
}
|
||||
|
||||
r.flow = append(r.flow, flowEntry{packet: fp})
|
||||
@@ -660,6 +693,10 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
||||
p.Release()
|
||||
return
|
||||
|
||||
case Drop:
|
||||
// Record it so the flow log shows the attempt, but never hand it to the receiver
|
||||
r.unlockedInjectFlow(sender, receiver, p, false)
|
||||
|
||||
case KeepRouting:
|
||||
fp := r.unlockedInjectFlow(sender, receiver, p, false)
|
||||
receiver.InjectUDPPacket(p)
|
||||
@@ -690,6 +727,85 @@ func (r *R) RouteUntilAfterMsgType(sender *nebula.Control, msgType header.Messag
|
||||
})
|
||||
}
|
||||
|
||||
// RouteFor routes everything that shows up for the given duration and then returns. Use it to let a test settle
|
||||
// deterministically rather than sleeping and hoping: a single FlushAll races a completing handshake, which queues
|
||||
// more packets right behind it.
|
||||
func (r *R) RouteFor(d time.Duration) {
|
||||
r.RouteForAllExitFuncOrTimeout(d, func(*udp.Packet, *nebula.Control) ExitType {
|
||||
return KeepRouting
|
||||
})
|
||||
}
|
||||
|
||||
// RouteForAllExitFuncOrTimeout is RouteForAllExitFunc with a deadline, reporting whether whatDo asked to exit
|
||||
// before time ran out. The unbounded version blocks forever on a quiet network, so this is what a test needs to
|
||||
// assert that something does NOT happen, or to route for a fixed settling period.
|
||||
func (r *R) RouteForAllExitFuncOrTimeout(timeout time.Duration, whatDo ExitFunc) bool {
|
||||
sc := make([]reflect.SelectCase, 0, len(r.controls)+1)
|
||||
cm := make([]*nebula.Control, 0, len(r.controls))
|
||||
|
||||
for _, c := range r.controls {
|
||||
sc = append(sc, reflect.SelectCase{
|
||||
Dir: reflect.SelectRecv,
|
||||
Chan: reflect.ValueOf(c.GetUDPTxChan()),
|
||||
Send: reflect.Value{},
|
||||
})
|
||||
cm = append(cm, c)
|
||||
}
|
||||
|
||||
timer := time.NewTimer(timeout)
|
||||
defer timer.Stop()
|
||||
sc = append(sc, reflect.SelectCase{
|
||||
Dir: reflect.SelectRecv,
|
||||
Chan: reflect.ValueOf(timer.C),
|
||||
Send: reflect.Value{},
|
||||
})
|
||||
|
||||
for {
|
||||
x, rx, _ := reflect.Select(sc)
|
||||
if x == len(cm) {
|
||||
return false
|
||||
}
|
||||
|
||||
r.Lock()
|
||||
p := rx.Interface().(*udp.Packet)
|
||||
receiver := r.getControl(cm[x].GetUDPAddr(), p.To, p)
|
||||
if receiver == nil {
|
||||
r.Unlock()
|
||||
panic("Can't RouteForAllExitFuncOrTimeout for host: " + p.To.String())
|
||||
}
|
||||
|
||||
e := whatDo(p, receiver)
|
||||
switch e {
|
||||
case ExitNow:
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
return true
|
||||
|
||||
case RouteAndExit:
|
||||
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
receiver.InjectUDPPacket(p)
|
||||
fp.WasReceived()
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
return true
|
||||
|
||||
case Drop:
|
||||
// Record it so the flow log shows the attempt, but never hand it to the receiver
|
||||
r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
|
||||
case KeepRouting:
|
||||
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
receiver.InjectUDPPacket(p)
|
||||
fp.WasReceived()
|
||||
|
||||
default:
|
||||
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
||||
}
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
}
|
||||
}
|
||||
|
||||
func (r *R) RouteForAllUntilAfterMsgTypeTo(receiver *nebula.Control, msgType header.MessageType, subType header.MessageSubType) {
|
||||
h := &header.H{}
|
||||
r.RouteForAllExitFunc(func(p *udp.Packet, r *nebula.Control) ExitType {
|
||||
@@ -782,6 +898,10 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
||||
p.Release()
|
||||
return
|
||||
|
||||
case Drop:
|
||||
// Record it so the flow log shows the attempt, but never hand it to the receiver
|
||||
r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
|
||||
case KeepRouting:
|
||||
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
receiver.InjectUDPPacket(p)
|
||||
|
||||
@@ -1,188 +0,0 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"log/slog"
|
||||
"testing"
|
||||
|
||||
"golang.org/x/net/ipv4"
|
||||
)
|
||||
|
||||
func TestInnerECN(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
pkt []byte
|
||||
want byte
|
||||
}{
|
||||
{"empty", nil, 0},
|
||||
{"v4_NotECT", v4WithToS(0x00), 0x00},
|
||||
{"v4_ECT0", v4WithToS(0x02), 0x02},
|
||||
{"v4_ECT1", v4WithToS(0x01), 0x01},
|
||||
{"v4_CE", v4WithToS(0x03), 0x03},
|
||||
{"v4_DSCP_then_NotECT", v4WithToS(0x88 | 0x00), 0x00},
|
||||
{"v4_DSCP_then_CE", v4WithToS(0x88 | 0x03), 0x03},
|
||||
{"v6_NotECT", v6WithTC(0x00), 0x00},
|
||||
{"v6_ECT0", v6WithTC(0x02), 0x02},
|
||||
{"v6_CE", v6WithTC(0x03), 0x03},
|
||||
{"v6_DSCP_then_CE", v6WithTC(0x88 | 0x03), 0x03},
|
||||
{"unknown_version", []byte{0xa5, 0xff}, 0},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
got := innerECN(c.pkt)
|
||||
if got != c.want {
|
||||
t.Errorf("innerECN=0x%02x want 0x%02x", got, c.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// v4WithToS returns a 2-byte slice tall enough for innerECN: byte 0 carries
|
||||
// version=4 in the high nibble, byte 1 is the full ToS so we exercise both
|
||||
// the DSCP and ECN portions through the byte 1 mask.
|
||||
func v4WithToS(tos byte) []byte {
|
||||
return []byte{0x45, tos}
|
||||
}
|
||||
|
||||
// v6WithTC builds a 2-byte slice that places a known traffic class value
|
||||
// across bytes 0 (high nibble of TC) and 1 (low nibble of TC). innerECN
|
||||
// extracts ECN as (b[1]>>4)&0x03, which corresponds to TC[1:0].
|
||||
func v6WithTC(tc byte) []byte {
|
||||
return []byte{0x60 | (tc>>4)&0x0f, (tc & 0x0f) << 4}
|
||||
}
|
||||
|
||||
func TestApplyOuterECN(t *testing.T) {
|
||||
silent := slog.New(slog.DiscardHandler)
|
||||
hi := &HostInfo{}
|
||||
|
||||
// Build a v4 packet helper with a given inner ECN field.
|
||||
v4 := func(innerECN byte) []byte {
|
||||
// 20-byte minimal IPv4 header with ToS = innerECN (DSCP zeroed).
|
||||
return []byte{
|
||||
0x45, innerECN, 0, 28,
|
||||
0, 0, 0x40, 0,
|
||||
64, 6, 0, 0,
|
||||
10, 0, 0, 1,
|
||||
10, 0, 0, 2,
|
||||
}
|
||||
}
|
||||
// Build a v6 packet helper with a given inner ECN field. ECN occupies
|
||||
// TC[1:0] which sit at byte 1 mask 0x30.
|
||||
v6 := func(innerECN byte) []byte {
|
||||
// 40-byte minimal IPv6 header with TC[1:0] = innerECN.
|
||||
pkt := make([]byte, 40)
|
||||
pkt[0] = 0x60 // version=6, TC[7:4]=0
|
||||
pkt[1] = (innerECN & 0x03) << 4 // TC[3:0]: low 2 bits = ECN, top 2 = DSCP-low (0)
|
||||
return pkt
|
||||
}
|
||||
|
||||
type cell struct {
|
||||
outer byte
|
||||
inner byte
|
||||
wantECN byte
|
||||
wantSame bool // expect inner unchanged (true => verify the byte didn't move)
|
||||
}
|
||||
|
||||
// RFC 6040 normal-mode combine table. Only outer==CE causes mutation.
|
||||
table := []cell{
|
||||
{ecnNotECT, ecnNotECT, ecnNotECT, true},
|
||||
{ecnNotECT, ecnECT0, ecnECT0, true},
|
||||
{ecnNotECT, ecnECT1, ecnECT1, true},
|
||||
{ecnNotECT, ecnCE, ecnCE, true},
|
||||
|
||||
{ecnECT0, ecnNotECT, ecnNotECT, true},
|
||||
{ecnECT0, ecnECT0, ecnECT0, true},
|
||||
{ecnECT0, ecnECT1, ecnECT1, true},
|
||||
{ecnECT0, ecnCE, ecnCE, true},
|
||||
|
||||
{ecnECT1, ecnNotECT, ecnNotECT, true},
|
||||
{ecnECT1, ecnECT0, ecnECT0, true},
|
||||
{ecnECT1, ecnECT1, ecnECT1, true},
|
||||
{ecnECT1, ecnCE, ecnCE, true},
|
||||
|
||||
{ecnCE, ecnNotECT, ecnNotECT, true}, // legacy: log, leave alone
|
||||
{ecnCE, ecnECT0, ecnCE, false}, // CE folded in
|
||||
{ecnCE, ecnECT1, ecnCE, false},
|
||||
{ecnCE, ecnCE, ecnCE, true},
|
||||
}
|
||||
|
||||
for _, c := range table {
|
||||
t.Run("v4", func(t *testing.T) {
|
||||
pkt := v4(c.inner)
|
||||
applyOuterECN(pkt, c.outer, hi, silent)
|
||||
got := pkt[1] & 0x03
|
||||
if got != c.wantECN {
|
||||
t.Errorf("v4 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
||||
}
|
||||
})
|
||||
t.Run("v6", func(t *testing.T) {
|
||||
pkt := v6(c.inner)
|
||||
applyOuterECN(pkt, c.outer, hi, silent)
|
||||
got := (pkt[1] >> 4) & 0x03
|
||||
if got != c.wantECN {
|
||||
t.Errorf("v6 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestApplyOuterECN_IPv4ChecksumStaysValid guards against H1: folding an outer
|
||||
// CE mark into the inner IPv4 ToS byte must keep the IPv4 header checksum valid.
|
||||
// The passthrough emit paths write the packet verbatim, so a stale checksum
|
||||
// turns an underlay congestion mark into packet loss at the receiver.
|
||||
func TestApplyOuterECN_IPv4ChecksumStaysValid(t *testing.T) {
|
||||
silent := slog.New(slog.DiscardHandler)
|
||||
hi := &HostInfo{}
|
||||
|
||||
// 20-byte IPv4 header with DSCP=0x88 and inner ECN = ECT(0). Folding CE
|
||||
// flips only the low two bits of the ToS byte while leaving DSCP intact.
|
||||
pkt := []byte{
|
||||
0x45, 0x88 | ecnECT0, 0, 40,
|
||||
0x1c, 0x46, 0x40, 0x00,
|
||||
64, 6, 0, 0,
|
||||
10, 0, 0, 1,
|
||||
10, 0, 0, 2,
|
||||
}
|
||||
// Stamp a correct header checksum before the fold.
|
||||
binary.BigEndian.PutUint16(pkt[10:12], ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]))
|
||||
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
|
||||
t.Fatal("test setup: initial header checksum invalid")
|
||||
}
|
||||
|
||||
applyOuterECN(pkt, ecnCE, hi, silent)
|
||||
|
||||
// CE folded in, DSCP preserved.
|
||||
if got, want := pkt[1], byte(0x88|ecnCE); got != want {
|
||||
t.Fatalf("ToS after fold = 0x%02x, want 0x%02x", got, want)
|
||||
}
|
||||
// The incremental RFC 1624 update must leave the checksum valid and equal
|
||||
// to a full recompute over the mutated header.
|
||||
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
|
||||
t.Fatalf("IPv4 header checksum invalid after CE fold: 0x%04x", binary.BigEndian.Uint16(pkt[10:12]))
|
||||
}
|
||||
if got, want := binary.BigEndian.Uint16(pkt[10:12]), ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]); got != want {
|
||||
t.Fatalf("checksum = 0x%04x, full recompute = 0x%04x", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// ipv4HeaderChecksum computes the RFC 1071 IPv4 header checksum over hdr,
|
||||
// treating the checksum field (bytes 10:12) as zero.
|
||||
func ipv4HeaderChecksum(hdr []byte) uint16 {
|
||||
var sum uint32
|
||||
for i := 0; i+1 < len(hdr); i += 2 {
|
||||
if i == 10 {
|
||||
continue // checksum field
|
||||
}
|
||||
sum += uint32(hdr[i])<<8 | uint32(hdr[i+1])
|
||||
}
|
||||
for sum > 0xffff {
|
||||
sum = (sum >> 16) + (sum & 0xffff)
|
||||
}
|
||||
return ^uint16(sum)
|
||||
}
|
||||
|
||||
// ipv4HeaderChecksumValid reports whether the stored checksum matches a fresh
|
||||
// computation over the header.
|
||||
func ipv4HeaderChecksumValid(hdr []byte) bool {
|
||||
return binary.BigEndian.Uint16(hdr[10:12]) == ipv4HeaderChecksum(hdr)
|
||||
}
|
||||
+29
-21
@@ -131,6 +131,9 @@ listen:
|
||||
port: 4242
|
||||
# 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
|
||||
# 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
|
||||
# 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)
|
||||
@@ -146,6 +149,14 @@ listen:
|
||||
# Default true; set to false to leave WDF in charge of inbound decisions on the listener port. Not reloadable.
|
||||
#windows_bypass_wdf: true
|
||||
|
||||
# On macOS only
|
||||
# macOS scopes the udp socket to the interface it was created on, so moving between networks (wifi to wired,
|
||||
# office to home) leaves Nebula sending out an interface that no longer has a route. When true, Nebula watches
|
||||
# the routing socket and rebinds the listener once the change settles.
|
||||
# iOS does not use this, the host app drives the same rebind itself.
|
||||
# Default true. Not reloadable.
|
||||
#rebind_on_network_change: true
|
||||
|
||||
# By default, Nebula replies to packets it has no tunnel for with a "recv_error" packet. This packet helps speed up reconnection
|
||||
# in the case that Nebula on either side did not shut down cleanly. This response can be abused as a way to discover if Nebula is running
|
||||
# on a host though. This option lets you configure if you want to send "recv_error" packets always, never, or only to private network remotes.
|
||||
@@ -161,6 +172,8 @@ listen:
|
||||
# allowing for more precise routing decisions based on the packet tags. Default is 0 meaning no mark is set.
|
||||
# This setting is reloadable.
|
||||
#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.
|
||||
# Currently, this defaults to 1 which means we have 1 tun queue reader and 1
|
||||
@@ -254,24 +267,29 @@ tun:
|
||||
# Default MTU for every packet, safe setting is (and the default) 1300 for internet based traffic
|
||||
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.
|
||||
#
|
||||
# When cpu_affinity is unset, nebula picks CPUs that do NOT service any physical NIC's interrupts (read from
|
||||
# /sys/class/net/*/device/msi_irqs and /proc/irq/*/effective_affinity_list): an encrypt thread pinned onto a core
|
||||
# that also runs NAPI for a NIC RX queue fights the softirq for the core and collapses throughput for flows hashed
|
||||
# to that queue. If the NIC's vectors blanket every allowed CPU (many drivers default to one queue per core) the
|
||||
# avoidance logs and falls back to the old spread; narrow the NIC's queue/IRQ spread (e.g. `ethtool -X <dev>
|
||||
# equal N`) or set cpu_affinity explicitly to benefit.
|
||||
# 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 and falls back to spreading queues across the allowed
|
||||
# CPUs. Setting this disables the automatic NIC-IRQ avoidance described under pin_threads — prefer CPUs that don't
|
||||
# service your underlay NIC's RX queue IRQs. Only meaningful while pin_threads is true. Not reloadable.
|
||||
# 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
|
||||
@@ -412,16 +430,6 @@ logging:
|
||||
# This setting is reloadable
|
||||
#inactivity_timeout: 10m
|
||||
|
||||
# ecn (default true) propagates ECN (Explicit Congestion Notification) across the tunnel per RFC 6040: the inner
|
||||
# packet's ECN codepoint is copied onto the outer carrier header on encapsulation, and an outer CE ("congestion
|
||||
# experienced") mark is folded back into the inner header on decapsulation. On linux it additionally stamps
|
||||
# RTAX_FEATURE_ECN on the routes nebula installs, so the kernel actively negotiates ECN for connections to mesh
|
||||
# prefixes. Disable this only when an underlay middlebox mangles or clears ECN bits unpredictably.
|
||||
# This setting is reloadable, BUT flipping it at runtime only updates the datapath (the inner<->outer copy/combine).
|
||||
# The RTAX_FEATURE_ECN flag on already-installed routes is NOT revisited on reload, so nebula must be restarted for
|
||||
# the route half of this setting to take effect.
|
||||
#ecn: true
|
||||
|
||||
# Nebula security group configuration
|
||||
firewall:
|
||||
# Action to take when a packet is not allowed by the firewall rules.
|
||||
|
||||
@@ -8,6 +8,15 @@ Before=sshd.service
|
||||
Type=notify
|
||||
NotifyAccess=main
|
||||
SyslogIdentifier=nebula
|
||||
|
||||
# Uncomment to run as an unprivileged user with only CAP_NET_ADMIN. Requires a
|
||||
# nebula user that owns the config directory. Add CAP_NET_BIND_SERVICE to both
|
||||
# lines if any listener (lighthouse DNS, listen.port, stats, sshd) binds <1024.
|
||||
#User=nebula
|
||||
#Group=nebula
|
||||
#CapabilityBoundingSet=CAP_NET_ADMIN
|
||||
#AmbientCapabilities=CAP_NET_ADMIN
|
||||
|
||||
ExecReload=/bin/kill -HUP $MAINPID
|
||||
ExecStart=/usr/local/bin/nebula -config /etc/nebula/config.yml
|
||||
Restart=always
|
||||
|
||||
+15
-14
@@ -21,6 +21,7 @@ import (
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
)
|
||||
|
||||
type FirewallInterface interface {
|
||||
@@ -262,11 +263,11 @@ func (f *Firewall) AddRule(incoming bool, proto uint8, startPort int32, endPort
|
||||
}
|
||||
|
||||
switch proto {
|
||||
case firewall.ProtoTCP:
|
||||
case iputil.IPProtocolTCP:
|
||||
fp = ft.TCP
|
||||
case firewall.ProtoUDP:
|
||||
case iputil.IPProtocolUDP:
|
||||
fp = ft.UDP
|
||||
case firewall.ProtoICMP, firewall.ProtoICMPv6:
|
||||
case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
|
||||
//ICMP traffic doesn't have ports, so we always coerce to "any", even if a value is provided
|
||||
if startPort != firewall.PortAny {
|
||||
f.l.Warn("ignoring port specification for ICMP firewall rule", "startPort", startPort)
|
||||
@@ -364,13 +365,13 @@ func AddFirewallRulesFromConfig(l *slog.Logger, inbound bool, c *config.C, fw Fi
|
||||
proto = firewall.ProtoAny
|
||||
startPort, endPort, err = parsePort(sPort)
|
||||
case "tcp":
|
||||
proto = firewall.ProtoTCP
|
||||
proto = iputil.IPProtocolTCP
|
||||
startPort, endPort, err = parsePort(sPort)
|
||||
case "udp":
|
||||
proto = firewall.ProtoUDP
|
||||
proto = iputil.IPProtocolUDP
|
||||
startPort, endPort, err = parsePort(sPort)
|
||||
case "icmp":
|
||||
proto = firewall.ProtoICMP
|
||||
proto = iputil.IPProtocolICMP
|
||||
startPort = firewall.PortAny
|
||||
endPort = firewall.PortAny
|
||||
if sPort != "" {
|
||||
@@ -560,9 +561,9 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
|
||||
}
|
||||
|
||||
switch fp.Protocol {
|
||||
case firewall.ProtoTCP:
|
||||
case iputil.IPProtocolTCP:
|
||||
c.Expires = time.Now().Add(f.TCPTimeout)
|
||||
case firewall.ProtoUDP:
|
||||
case iputil.IPProtocolUDP:
|
||||
c.Expires = time.Now().Add(f.UDPTimeout)
|
||||
default:
|
||||
c.Expires = time.Now().Add(f.DefaultTimeout)
|
||||
@@ -582,9 +583,9 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
|
||||
c := &conn{}
|
||||
|
||||
switch fp.Protocol {
|
||||
case firewall.ProtoTCP:
|
||||
case iputil.IPProtocolTCP:
|
||||
timeout = f.TCPTimeout
|
||||
case firewall.ProtoUDP:
|
||||
case iputil.IPProtocolUDP:
|
||||
timeout = f.UDPTimeout
|
||||
default:
|
||||
timeout = f.DefaultTimeout
|
||||
@@ -635,15 +636,15 @@ func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedC
|
||||
}
|
||||
|
||||
switch p.Protocol {
|
||||
case firewall.ProtoTCP:
|
||||
case iputil.IPProtocolTCP:
|
||||
if ft.TCP.match(p, incoming, c, caPool) {
|
||||
return true
|
||||
}
|
||||
case firewall.ProtoUDP:
|
||||
case iputil.IPProtocolUDP:
|
||||
if ft.UDP.match(p, incoming, c, caPool) {
|
||||
return true
|
||||
}
|
||||
case firewall.ProtoICMP, firewall.ProtoICMPv6:
|
||||
case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
|
||||
if ft.ICMP.match(p, incoming, c, caPool) {
|
||||
return true
|
||||
}
|
||||
@@ -680,7 +681,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
|
||||
if p.Protocol == firewall.ProtoICMP || p.Protocol == firewall.ProtoICMPv6 {
|
||||
if p.Protocol == iputil.IPProtocolICMP || p.Protocol == iputil.IPProtocolICMPv6 {
|
||||
// port numbers are re-used for connection tracking of ICMP,
|
||||
// but we don't want to actually filter on them.
|
||||
return fp[firewall.PortAny].match(p, c, caPool)
|
||||
|
||||
+16
-10
@@ -4,17 +4,14 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
)
|
||||
|
||||
type m = map[string]any
|
||||
|
||||
const (
|
||||
ProtoAny = 0 // When we want to handle HOPOPT (0) we can change this, if ever
|
||||
ProtoTCP = 6
|
||||
ProtoUDP = 17
|
||||
ProtoICMP = 1
|
||||
ProtoICMPv6 = 58
|
||||
|
||||
ProtoAny = 0 // When we want to handle HOPOPT (0) we can change this, if ever
|
||||
PortAny = 0 // Special value for matching `port: any`
|
||||
PortFragment = -1 // Special value for matching `port: fragment`
|
||||
)
|
||||
@@ -45,13 +42,13 @@ func (fp *Packet) Copy() *Packet {
|
||||
func (fp Packet) MarshalJSON() ([]byte, error) {
|
||||
var proto string
|
||||
switch fp.Protocol {
|
||||
case ProtoTCP:
|
||||
case iputil.IPProtocolTCP:
|
||||
proto = "tcp"
|
||||
case ProtoICMP:
|
||||
case iputil.IPProtocolICMP:
|
||||
proto = "icmp"
|
||||
case ProtoICMPv6:
|
||||
case iputil.IPProtocolICMPv6:
|
||||
proto = "icmpv6"
|
||||
case ProtoUDP:
|
||||
case iputil.IPProtocolUDP:
|
||||
proto = "udp"
|
||||
default:
|
||||
proto = fmt.Sprintf("unknown %v", fp.Protocol)
|
||||
@@ -65,3 +62,12 @@ func (fp Packet) MarshalJSON() ([]byte, error) {
|
||||
"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
|
||||
}
|
||||
|
||||
+34
-33
@@ -13,6 +13,7 @@ import (
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -72,20 +73,20 @@ func TestFirewall_AddRule(t *testing.T) {
|
||||
ti6, err := netip.ParsePrefix("fd12::34/128")
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoTCP, 1, 1, []string{}, "", "", "", "", ""))
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolTCP, 1, 1, []string{}, "", "", "", "", ""))
|
||||
// An empty rule is 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.Hosts)
|
||||
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", ""))
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", ""))
|
||||
assert.Nil(t, fw.InRules.UDP[1].Any.Any)
|
||||
assert.Contains(t, fw.InRules.UDP[1].Any.Groups[0].Groups, "g1")
|
||||
assert.Empty(t, fw.InRules.UDP[1].Any.Hosts)
|
||||
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoICMP, 1, 1, []string{}, "h1", "", "", "", ""))
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 1, 1, []string{}, "h1", "", "", "", ""))
|
||||
//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.Empty(t, fw.InRules.ICMP[firewall.PortAny].Any.Groups)
|
||||
@@ -116,11 +117,11 @@ func TestFirewall_AddRule(t *testing.T) {
|
||||
assert.True(t, ok)
|
||||
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", ""))
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", ""))
|
||||
assert.Contains(t, fw.InRules.UDP[1].CANames, "ca-name")
|
||||
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha"))
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha"))
|
||||
assert.Contains(t, fw.InRules.UDP[1].CAShas, "ca-sha")
|
||||
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||
@@ -185,7 +186,7 @@ func TestFirewall_Drop(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
LocalPort: 10,
|
||||
RemotePort: 90,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
@@ -263,7 +264,7 @@ func TestFirewall_DropV6(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("fd12::34"),
|
||||
LocalPort: 10,
|
||||
RemotePort: 90,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
@@ -350,7 +351,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
Certificate: &dummyCert{},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoUDP}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolUDP}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -360,7 +361,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
Certificate: &dummyCert{},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 1}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 1}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -370,7 +371,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
}
|
||||
ip := netip.MustParsePrefix("9.254.254.254/32")
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
||||
}
|
||||
})
|
||||
b.Run("pass proto, port, fail on local CIDRv6", func(b *testing.B) {
|
||||
@@ -379,7 +380,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
}
|
||||
ip := netip.MustParsePrefix("fd99::99/128")
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -392,7 +393,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
|
||||
}
|
||||
})
|
||||
b.Run("pass proto, port, any local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
|
||||
@@ -404,7 +405,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -417,7 +418,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, 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) {
|
||||
@@ -429,7 +430,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -441,7 +442,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"good-group": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -453,7 +454,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"good-group": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
|
||||
}
|
||||
})
|
||||
b.Run("pass on group on specific local cidr6", func(b *testing.B) {
|
||||
@@ -464,7 +465,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"good-group": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -476,7 +477,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp)
|
||||
ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -492,7 +493,7 @@ func TestFirewall_Drop2(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
LocalPort: 10,
|
||||
RemotePort: 90,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
@@ -550,7 +551,7 @@ func TestFirewall_Drop3(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
LocalPort: 1,
|
||||
RemotePort: 1,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
@@ -638,7 +639,7 @@ func TestFirewall_Drop3V6(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("fd12::34"),
|
||||
LocalPort: 1,
|
||||
RemotePort: 1,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
@@ -675,7 +676,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
LocalPort: 10,
|
||||
RemotePort: 90,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
network := netip.MustParsePrefix("1.2.3.4/24")
|
||||
@@ -758,13 +759,13 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||
templ := firewall.Packet{
|
||||
LocalAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
Protocol: firewall.ProtoICMP,
|
||||
Protocol: iputil.IPProtocolICMP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
t.Run("ICMP allowed", func(t *testing.T) {
|
||||
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoICMP, 0, 0, []string{"any"}, "", "", "", "", ""))
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 0, 0, []string{"any"}, "", "", "", "", ""))
|
||||
t.Run("zero ports", func(t *testing.T) {
|
||||
p := templ.Copy()
|
||||
p.LocalPort = 0
|
||||
@@ -910,7 +911,7 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("192.0.2.3"),
|
||||
LocalPort: 1,
|
||||
RemotePort: 1,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
|
||||
@@ -961,7 +962,7 @@ func TestFirewall_ConntrackSourceSpoofingAcrossPeers(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
||||
LocalPort: 443,
|
||||
RemotePort: 55000,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
}
|
||||
|
||||
require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil),
|
||||
@@ -1031,7 +1032,7 @@ func BenchmarkFirewallDropConntrackHit(b *testing.B) {
|
||||
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
||||
LocalPort: 443,
|
||||
RemotePort: 55000,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
@@ -1317,28 +1318,28 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
|
||||
mf := &mockFirewall{}
|
||||
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))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding udp rule
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
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))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding icmp rule
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
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))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding icmp rule no port
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"proto": "icmp", "host": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding any rule
|
||||
conf = config.NewC(test.NewLogger())
|
||||
@@ -1582,7 +1583,7 @@ func buildTestCase(setup testsetup, err error, theirPrefixes ...netip.Prefix) te
|
||||
RemoteAddr: theirPrefixes[0].Addr(),
|
||||
LocalPort: 10,
|
||||
RemotePort: 90,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
return testcase{
|
||||
|
||||
@@ -1,13 +1,12 @@
|
||||
module github.com/slackhq/nebula
|
||||
|
||||
go 1.25.0
|
||||
go 1.26.0
|
||||
|
||||
require (
|
||||
dario.cat/mergo v1.0.2
|
||||
filippo.io/bigmod v0.1.0
|
||||
github.com/anmitsu/go-shlex v0.0.0-20200514113438-38f4b401e2be
|
||||
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/gaissmai/bart v0.28.0
|
||||
github.com/gogo/protobuf v1.3.2
|
||||
@@ -16,20 +15,19 @@ require (
|
||||
github.com/miekg/dns v1.1.72
|
||||
github.com/miekg/pkcs11 v1.1.2
|
||||
github.com/nbrownus/go-metrics-prometheus v0.0.0-20210712211119-974a6260965f
|
||||
github.com/prometheus/client_golang v1.23.2
|
||||
github.com/prometheus/client_golang v1.24.1
|
||||
github.com/rcrowley/go-metrics v0.0.0-20201227073835-cf1acfcdf475
|
||||
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e
|
||||
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
||||
github.com/stretchr/testify v1.11.1
|
||||
github.com/stretchr/testify v1.12.0
|
||||
github.com/vishvananda/netlink v1.3.1
|
||||
go.uber.org/goleak v1.3.0
|
||||
go.yaml.in/yaml/v3 v3.0.4
|
||||
golang.org/x/crypto v0.53.0
|
||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
||||
golang.org/x/net v0.56.0
|
||||
golang.org/x/sync v0.21.0
|
||||
golang.org/x/sys v0.46.0
|
||||
golang.org/x/term v0.44.0
|
||||
go.yaml.in/yaml/v3 v3.0.5
|
||||
golang.org/x/crypto v0.56.0
|
||||
golang.org/x/net v0.58.0
|
||||
golang.org/x/sync v0.22.0
|
||||
golang.org/x/sys v0.47.0
|
||||
golang.org/x/term v0.45.0
|
||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
|
||||
golang.zx2c4.com/wireguard/windows v1.0.1
|
||||
@@ -41,15 +39,12 @@ require (
|
||||
require (
|
||||
github.com/beorn7/perks v1.0.1 // 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/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/common v0.66.1 // indirect
|
||||
github.com/prometheus/procfs v0.16.1 // indirect
|
||||
github.com/prometheus/common v0.70.1 // indirect
|
||||
github.com/prometheus/procfs v0.21.1 // indirect
|
||||
github.com/vishvananda/netns v0.0.5 // indirect
|
||||
go.yaml.in/yaml/v2 v2.4.2 // indirect
|
||||
golang.org/x/mod v0.36.0 // indirect
|
||||
golang.org/x/time v0.5.0 // indirect
|
||||
golang.org/x/tools v0.45.0 // indirect
|
||||
|
||||
@@ -19,10 +19,7 @@ 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.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
||||
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.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||
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/go.mod h1:xbMo+0i6+IGbYdJhF31t2eR1BIU0CYc12+BNAKwUTag=
|
||||
@@ -70,15 +67,14 @@ github.com/kardianos/service v1.3.0 h1:/LGy+xPP2TM+GLTiCZ2di7cy0Jd/qrawlTUfqKYFd
|
||||
github.com/kardianos/service v1.3.0/go.mod h1:E4V9ufUuY82F7Ztlu1eN9VXWIQxg8NoLQlmFe0MtrXc=
|
||||
github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8=
|
||||
github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck=
|
||||
github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo=
|
||||
github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYWRCY2AiWywWQ=
|
||||
github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk=
|
||||
github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
||||
github.com/konsorten/go-windows-terminal-sequences v1.0.1/go.mod h1:T0+1ngSBFLxvqU3pZ+m/2kptfBszLMUkC4ZK/EgS/cQ=
|
||||
github.com/konsorten/go-windows-terminal-sequences v1.0.3/go.mod h1:T0+1ngSBFLxvqU3pZ+m/2kptfBszLMUkC4ZK/EgS/cQ=
|
||||
github.com/kr/logfmt v0.0.0-20140226030751-b84e30acd515/go.mod h1:+0opPa2QZZtGFBFZlji/RkVcI2GknAs/DXo4wKdlNEc=
|
||||
github.com/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORNo=
|
||||
github.com/kr/pretty v0.2.1 h1:Fmg33tUaq4/8ym9TJN1x7sLJnHVwhP33CNkpYV/7rwI=
|
||||
github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI=
|
||||
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
||||
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
||||
github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ=
|
||||
github.com/kr/text v0.1.0 h1:45sCR5RtlFHMR4UwH9sdQ5TC8v0qDQCHnXt+kaKSTVE=
|
||||
github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI=
|
||||
@@ -102,14 +98,13 @@ 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.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/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.7.1/go.mod h1:PY5Wy2awLA44sXw4AOSfFBetzPP4j5+D6mVACh+pe2M=
|
||||
github.com/prometheus/client_golang v1.11.0/go.mod h1:Z6t4BnS23TR94PD6BsDNk8yVqroYurpAkEiz0P2BEV0=
|
||||
github.com/prometheus/client_golang v1.23.2 h1:Je96obch5RDVy3FDMndoUsjAhG5Edi49h0RJWRi/o0o=
|
||||
github.com/prometheus/client_golang v1.23.2/go.mod h1:Tb1a6LWHB3/SPIzCoaDXI4I8UHKeFTEQ1YCr+0Gyqmg=
|
||||
github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU=
|
||||
github.com/prometheus/client_golang v1.24.1/go.mod h1:F+oSRECHg4sse5ucfYpYDeIv/hu68Zo0uoHKetWnzcE=
|
||||
github.com/prometheus/client_model v0.0.0-20180712105110-5c3871d89910/go.mod h1:MbSGuTsp3dbXC40dX6PRTWyKYBIrTGTE9sqQNg2J8bo=
|
||||
github.com/prometheus/client_model v0.0.0-20190129233127-fd36f4220a90/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA=
|
||||
github.com/prometheus/client_model v0.2.0/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA=
|
||||
@@ -118,18 +113,16 @@ github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvM
|
||||
github.com/prometheus/common v0.4.1/go.mod h1:TNfzLD0ON7rHzMJeJkieUDPYmFC7Snx/y86RQel1bk4=
|
||||
github.com/prometheus/common v0.10.0/go.mod h1:Tlit/dnDKsSWFlCLTWaA1cyBgKHSMdTB80sz/V91rCo=
|
||||
github.com/prometheus/common v0.26.0/go.mod h1:M7rCNAaPfAosfx8veZJCuw84e35h3Cfd9VFqTh1DIvc=
|
||||
github.com/prometheus/common v0.66.1 h1:h5E0h5/Y8niHc5DlaLlWLArTQI7tMrsfQjHV+d9ZoGs=
|
||||
github.com/prometheus/common v0.66.1/go.mod h1:gcaUsgf3KfRSwHY4dIMXLPV0K/Wg1oZ8+SbZk/HH/dA=
|
||||
github.com/prometheus/common v0.70.1 h1:1HvjP4D5oL3t8RsPlwxA9onvvStjtIHYE5XuuwOi/PY=
|
||||
github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc=
|
||||
github.com/prometheus/procfs v0.0.0-20181005140218-185b4288413d/go.mod h1:c3At6R/oaqEKCNdg8wHV1ftS6bRYblBhIjjI8uT2IGk=
|
||||
github.com/prometheus/procfs v0.0.2/go.mod h1:TjEm7ze935MbeOT/UhFTIMYKhuLP4wbCsTZCD3I8kEA=
|
||||
github.com/prometheus/procfs v0.1.3/go.mod h1:lV6e/gmhEcM9IjHGsFOCxxuZ+z1YqCvr4OA4YeYWdaU=
|
||||
github.com/prometheus/procfs v0.6.0/go.mod h1:cz+aTbrPOrUb4q7XlbU9ygM+/jj0fzG6c1xBZuNvfVA=
|
||||
github.com/prometheus/procfs v0.16.1 h1:hZ15bTNuirocR6u0JZ6BAHHmwS1p8B4P6MRqxtzMyRg=
|
||||
github.com/prometheus/procfs v0.16.1/go.mod h1:teAbpZRB1iIAJYREa1LsoWUXykVXA1KlTmWl8x/U+Is=
|
||||
github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI=
|
||||
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
|
||||
github.com/rcrowley/go-metrics v0.0.0-20201227073835-cf1acfcdf475 h1:N/ElC8H3+5XpJzTSTfLsJV/mx9Q9g7kxmchpfZyxgzM=
|
||||
github.com/rcrowley/go-metrics v0.0.0-20201227073835-cf1acfcdf475/go.mod h1:bCqnVzQkZxMG4s8nGwiZ5l3QUCyqpo9Y+/ZMZ9VjZe4=
|
||||
github.com/rogpeppe/go-internal v1.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjRBZyWFQ=
|
||||
github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncjaFoBhdsK/akog=
|
||||
github.com/sirupsen/logrus v1.2.0/go.mod h1:LxeOpSwHxABJmUn/MG1IvRgCAasNZTLOkJPxbbu5VWo=
|
||||
github.com/sirupsen/logrus v1.4.2/go.mod h1:tLMulIdttU9McNUspp0xgXVQah82FyeX6MwdIuYE2rE=
|
||||
github.com/sirupsen/logrus v1.6.0/go.mod h1:7uNnSEd1DgxDLC74fIahvMZmmYsHGZGEOFrfsX/uA88=
|
||||
@@ -143,8 +136,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.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.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
||||
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||
github.com/stretchr/testify v1.12.0 h1:K6Mr6jO9JICuend/5xzTM03ydSV3vdNRYAdPSukj8uI=
|
||||
github.com/stretchr/testify v1.12.0/go.mod h1:bOYBZb5qJ00vPzWfIqBUZPaxK8jWiXc6d3ErP4Ca9Gw=
|
||||
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/netns v0.0.5 h1:DfiHV+j8bA32MFM7bfEunvT8IAqQ/NzSJHtcmW5zdEY=
|
||||
@@ -153,19 +146,17 @@ github.com/yuin/goldmark v1.1.27/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9de
|
||||
github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74=
|
||||
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||
go.yaml.in/yaml/v2 v2.4.2 h1:DzmwEr2rDGHl7lsFgAHxmNz/1NlQ7xLIrlN2h5d1eGI=
|
||||
go.yaml.in/yaml/v2 v2.4.2/go.mod h1:081UH+NErpNdqlCXm3TtEran0rJZGxAYx9hb/ELlsPU=
|
||||
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
|
||||
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
||||
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
|
||||
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
|
||||
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
||||
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg=
|
||||
golang.org/x/crypto v0.0.0-20180904163835-0709b304e793/go.mod h1:6SG95UA2DQfeDnfUPMdvaQW0Q7yPrPDi9nlGo2tz2b4=
|
||||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
||||
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
|
||||
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
|
||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
|
||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
|
||||
golang.org/x/crypto v0.56.0 h1:GUh5Ii4J5jtcseSMiRqr1jXCNHoxjeV9Fmekc2oLy6Y=
|
||||
golang.org/x/crypto v0.56.0/go.mod h1:OMW5y6CY9l38uPLmxU6l6pwcXp1obtLo3e6gT7gQR2I=
|
||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||
golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg=
|
||||
golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||
@@ -182,8 +173,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
|
||||
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
||||
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
||||
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
|
||||
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
|
||||
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
|
||||
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
@@ -191,8 +182,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
||||
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=
|
||||
golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
||||
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
@@ -208,11 +199,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w
|
||||
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
|
||||
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||
golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
|
||||
golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y=
|
||||
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
|
||||
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
|
||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
|
||||
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
|
||||
+117
@@ -0,0 +1,117 @@
|
||||
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
|
||||
}
|
||||
+29
-7
@@ -295,7 +295,13 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
||||
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
|
||||
err := hm.outside.WriteTo(stage0, addr)
|
||||
if err != nil {
|
||||
hostinfo.logger(hm.l).Error("Failed to send handshake message",
|
||||
// These repeat every attempt, so match the success log below and only shout when the remotes changed
|
||||
level := slog.LevelDebug
|
||||
if remotesHaveChanged {
|
||||
level = slog.LevelError
|
||||
}
|
||||
|
||||
hostinfo.logger(hm.l).Log(context.Background(), level, "Failed to send handshake message",
|
||||
"udpAddr", addr,
|
||||
"initiatorIndex", hostinfo.localIndexId,
|
||||
"handshake", hsFields,
|
||||
@@ -529,7 +535,9 @@ func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
||||
|
||||
func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
||||
for _, addr := range hostinfo.vpnAddrs {
|
||||
delete(hm.vpnIps, addr)
|
||||
if cur, ok := hm.vpnIps[addr]; ok && cur.hostinfo == hostinfo {
|
||||
delete(hm.vpnIps, addr)
|
||||
}
|
||||
}
|
||||
|
||||
if len(hm.vpnIps) == 0 {
|
||||
@@ -741,8 +749,14 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head
|
||||
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{
|
||||
ConnectionState: newConnectionStateFromResult(result),
|
||||
ConnectionState: connState,
|
||||
localIndexId: result.LocalIndex,
|
||||
remoteIndexId: result.RemoteIndex,
|
||||
vpnAddrs: vpnAddrs,
|
||||
@@ -860,7 +874,13 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
||||
}
|
||||
|
||||
// Handshake complete; build the ConnectionState now that we have keys and a verified peer cert.
|
||||
hostinfo.ConnectionState = newConnectionStateFromResult(result)
|
||||
cs, err := 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
|
||||
if remoteCert == nil {
|
||||
@@ -967,7 +987,9 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
||||
nb := make([]byte, 12, 12)
|
||||
out := make([]byte, mtu)
|
||||
for _, cp := range hh.packetStore {
|
||||
//todo use a sendbatcher
|
||||
// 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)
|
||||
}
|
||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||
@@ -1078,8 +1100,8 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos
|
||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||
// We received a valid handshake on this relay, so make sure the relay
|
||||
// state reflects that, in case it had been marked Disestablished.
|
||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
|
||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established)
|
||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false, 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
|
||||
}
|
||||
|
||||
func (mw *mockEncWriter) SendVia(_ *HostInfo, _ *Relay, _, _, _ []byte, _ bool) {
|
||||
func (mw *mockEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
+11
-6
@@ -190,13 +190,18 @@ func SubTypeName(t MessageType, s MessageSubType) string {
|
||||
}
|
||||
|
||||
func IsValidSubType(t MessageType, s MessageSubType) bool {
|
||||
if n, ok := subTypeMap[t]; ok {
|
||||
if _, ok := (*n)[s]; ok {
|
||||
return true
|
||||
}
|
||||
switch t {
|
||||
case Message:
|
||||
return s == MessageNone || s == MessageRelay
|
||||
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
|
||||
|
||||
@@ -102,6 +102,57 @@ func TestTypeMap(t *testing.T) {
|
||||
}, 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) {
|
||||
assert.Equal(
|
||||
t,
|
||||
|
||||
+78
-14
@@ -239,11 +239,15 @@ const (
|
||||
|
||||
type HostInfo struct {
|
||||
remote atomic.Pointer[netip.AddrPort]
|
||||
remotes *RemoteList
|
||||
promoteCounter atomic.Uint32
|
||||
ConnectionState *ConnectionState
|
||||
remoteIndexId uint32
|
||||
localIndexId uint32
|
||||
|
||||
// Traffic bits, pendingDeletion, and the rebind epoch we last sent under
|
||||
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
|
||||
// The host may have other vpn addresses that are outside our
|
||||
@@ -262,11 +266,6 @@ type HostInfo struct {
|
||||
// This is used to limit lighthouse re-queries in chatty clients
|
||||
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
|
||||
// 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
|
||||
@@ -275,9 +274,6 @@ type HostInfo struct {
|
||||
lastRoam time.Time
|
||||
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.
|
||||
// This value will be behind against actual tunnel utilization in the hot path.
|
||||
// This should only be used by the ConnectionManagers ticker routine.
|
||||
@@ -287,7 +283,6 @@ type HostInfo struct {
|
||||
type ViaSender struct {
|
||||
UdpAddr netip.AddrPort
|
||||
relayHI *HostInfo // relayHI is the host info object of the relay
|
||||
remoteIdx uint32 // remoteIdx is the index included in the header of the received packet
|
||||
relay *Relay // relay contains the rest of the relay information, including the PeerIP of the host trying to communicate with us.
|
||||
IsRelayed bool // IsRelayed is true if the packet was sent through a relay
|
||||
}
|
||||
@@ -544,6 +539,17 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
||||
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 {
|
||||
hm.RLock()
|
||||
if h, ok := hm.Indexes[index]; ok {
|
||||
@@ -659,7 +665,7 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
|
||||
hm.Indexes[hostinfo.localIndexId] = hostinfo
|
||||
hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo
|
||||
|
||||
hostinfo.out.Store(true)
|
||||
hostinfo.markOut(f.rebindEpoch.Load())
|
||||
if f.connectionManager != nil { // f.connectionManager is only nil in some unit tests
|
||||
f.connectionManager.trafficTimer.Add(hostinfo.localIndexId, f.connectionManager.checkInterval)
|
||||
}
|
||||
@@ -760,6 +766,64 @@ 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 {
|
||||
if i.ConnectionState != nil {
|
||||
return i.ConnectionState.peerCert
|
||||
|
||||
@@ -401,3 +401,49 @@ func TestHostMap_RelayState(t *testing.T) {
|
||||
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,6 +2,7 @@ package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
@@ -15,7 +16,7 @@ import (
|
||||
"github.com/slackhq/nebula/routing"
|
||||
)
|
||||
|
||||
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.ParsedPacket, nb []byte, sendBatch *batch.SendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
|
||||
// only valid until the next Read on that queue. Every consumer below
|
||||
// (parse, self-forward, handshake cache, sendInsideMessage) reads it
|
||||
@@ -57,6 +58,9 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
|
||||
// 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
|
||||
})
|
||||
@@ -74,7 +78,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
|
||||
return
|
||||
}
|
||||
|
||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(&fwPacket.Packet, func(hh *HandshakeHostInfo) {
|
||||
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
||||
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
||||
// so retaining segments past the loop is safe.
|
||||
@@ -105,9 +109,9 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
|
||||
return
|
||||
}
|
||||
|
||||
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||
dropReason := f.firewall.Drop(fwPacket.Packet, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||
if dropReason == nil {
|
||||
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch, rejectBuf, q)
|
||||
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch)
|
||||
} else {
|
||||
f.rejectInside(packet, rejectBuf, q)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
@@ -126,7 +130,6 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
|
||||
c := ci.messageCounter.Add(1)
|
||||
|
||||
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
||||
f.connectionManager.Out(hostinfo)
|
||||
|
||||
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
|
||||
if noiseutil.EncryptLockNeeded {
|
||||
@@ -138,8 +141,7 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
|
||||
"udpAddr", hostinfo.GetRemote(),
|
||||
"counter", c,
|
||||
)
|
||||
// Skip this segment; the rest of the superpacket can still
|
||||
// go out — TCP will retransmit anything we drop here.
|
||||
// Skip this segment; the rest of the superpacket can still go out. TCP will retransmit anything we drop here.
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -151,28 +153,28 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
|
||||
// later sendmmsg flush. Segmentation is fused with encryption here so the
|
||||
// kernel-supplied superpacket bytes never get written into a separate
|
||||
// scratch arena: SegmentSuperpacket builds each segment's plaintext in
|
||||
// segScratch[:segLen] in turn, and we encrypt directly into a fresh
|
||||
// SendBatch slot.
|
||||
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int) {
|
||||
// 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
|
||||
}
|
||||
|
||||
remote := hostinfo.GetRemote()
|
||||
ecnEnabled := f.ecnEnabled.Load()
|
||||
if hostinfo.lastRebindCount != f.rebindCount {
|
||||
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
||||
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||
// 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])
|
||||
hostinfo.lastRebindCount = f.rebindCount
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind counter",
|
||||
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
|
||||
@@ -211,11 +213,7 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
|
||||
return nil
|
||||
}
|
||||
|
||||
var ecn byte
|
||||
if ecnEnabled {
|
||||
ecn = innerECN(seg)
|
||||
}
|
||||
sendBatch.Commit(toSend, relayHostInfo.GetRemote(), ecn)
|
||||
sendBatch.Commit(toSend, relayHostInfo.GetRemote())
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
@@ -233,36 +231,14 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
|
||||
return nil
|
||||
}
|
||||
|
||||
var ecn byte
|
||||
if ecnEnabled {
|
||||
ecn = innerECN(seg)
|
||||
}
|
||||
sendBatch.Commit(out, remote, ecn)
|
||||
sendBatch.Commit(out, remote)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Error("Failed to segment superpacket for send",
|
||||
"error", err,
|
||||
)
|
||||
hostinfo.logger(f.l).Error("Failed to segment superpacket for send", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// innerECN returns the 2-bit IP-level ECN codepoint of an inner IPv4 or IPv6
|
||||
// packet, or 0 if pkt is too short or its IP version is unrecognized. Used at
|
||||
// encap to copy the inner codepoint onto the outer carrier per RFC 6040.
|
||||
func innerECN(pkt []byte) byte {
|
||||
if len(pkt) < 2 {
|
||||
return 0
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
return pkt[1] & 0x03
|
||||
case 6:
|
||||
return (pkt[1] >> 4) & 0x03
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||
if !f.firewall.OutboundSendReject {
|
||||
return
|
||||
@@ -279,27 +255,30 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
|
||||
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, rejectBuf []byte, q int) {
|
||||
if !f.firewall.InboundSendReject {
|
||||
return
|
||||
}
|
||||
|
||||
out = iputil.CreateRejectPacket(packet, out)
|
||||
// split rejectBuf to make sure we have room to write the plaintext rejection, then encrypt it, without trampling anything
|
||||
// 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 {
|
||||
return
|
||||
}
|
||||
|
||||
if len(out) > iputil.MaxRejectPacketSize {
|
||||
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
|
||||
}
|
||||
|
||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, packet, q)
|
||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, encryptBuf, 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
|
||||
@@ -393,7 +372,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) {
|
||||
fp := &firewall.Packet{}
|
||||
fp := &firewall.ParsedPacket{}
|
||||
err := newPacket(p, false, fp)
|
||||
if err != nil {
|
||||
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
|
||||
@@ -401,7 +380,7 @@ func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubTyp
|
||||
}
|
||||
|
||||
// check if packet is in outbound fw rules
|
||||
dropReason := f.firewall.Drop(*fp, false, hostinfo, f.pki.GetCAPool(), nil)
|
||||
dropReason := f.firewall.Drop(fp.Packet, false, hostinfo, f.pki.GetCAPool(), nil)
|
||||
if dropReason != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("dropping cached packet",
|
||||
@@ -452,6 +431,14 @@ func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *C
|
||||
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.
|
||||
func (f *Interface) dropExhausted(hostinfo *HostInfo, c uint64, msg string) {
|
||||
f.messageMetrics.TxExhausted(1)
|
||||
if c == RejectAfterMessages {
|
||||
hostinfo.logger(f.l).Error(msg)
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) prepareSendVia(via *HostInfo,
|
||||
relay *Relay,
|
||||
ad,
|
||||
@@ -463,10 +450,17 @@ func (f *Interface) prepareSendVia(via *HostInfo,
|
||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||
via.ConnectionState.writeLock.Lock()
|
||||
}
|
||||
c := via.ConnectionState.messageCounter.Add(1)
|
||||
c, ok := via.ConnectionState.NextMessageCounter()
|
||||
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)
|
||||
f.connectionManager.Out(via)
|
||||
f.connectionManager.OutNoRebind(via)
|
||||
|
||||
// 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.
|
||||
@@ -514,20 +508,14 @@ func (f *Interface) prepareSendVia(via *HostInfo,
|
||||
// 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,
|
||||
) {
|
||||
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
|
||||
}
|
||||
|
||||
err = f.writers[0].WriteTo(toSend, via.GetRemote())
|
||||
err = f.writers[q].WriteTo(toSend, via.GetRemote())
|
||||
if err != nil {
|
||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
||||
}
|
||||
@@ -554,21 +542,24 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||
ci.writeLock.Lock()
|
||||
}
|
||||
c := ci.messageCounter.Add(1)
|
||||
c, ok := ci.NextMessageCounter()
|
||||
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)
|
||||
out = header.Encode(out, header.Version, t, st, hostinfo.remoteIndexId, c)
|
||||
f.connectionManager.Out(hostinfo)
|
||||
|
||||
// Query our LH if we haven't since the last time we've been rebound, this will cause the remote to punch against
|
||||
// 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.
|
||||
// A closing tunnel is torn down right after this, so skip the connection manager entirely: no point recording
|
||||
// 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.
|
||||
if t != header.CloseTunnel && f.connectionManager.Out(hostinfo) {
|
||||
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
|
||||
hostinfo.lastRebindCount = f.rebindCount
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("Lighthouse update triggered for punch due to rebind counter",
|
||||
f.l.Debug("Lighthouse update triggered for punch due to rebind epoch",
|
||||
"vpnAddrs", hostinfo.vpnAddrs,
|
||||
)
|
||||
}
|
||||
@@ -601,7 +592,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
||||
"error", err,
|
||||
"udpAddr", remote,
|
||||
"udpAddr", hr,
|
||||
)
|
||||
}
|
||||
} else {
|
||||
@@ -616,7 +607,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
)
|
||||
continue
|
||||
}
|
||||
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true)
|
||||
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true, q)
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
+265
@@ -0,0 +1,265 @@
|
||||
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:]))
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
+95
-81
@@ -2,6 +2,7 @@ package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/fips140"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
@@ -96,12 +97,7 @@ type Interface struct {
|
||||
// pinThreads controls whether listenIn pins each TUN reader OS thread to
|
||||
// a CPU at all (tun.pin_threads, default true). When false, threads are
|
||||
// left free to migrate as on stock nebula.
|
||||
pinThreads bool
|
||||
// ecnEnabled gates RFC 6040 underlay ECN propagation. When true,
|
||||
// inside.go copies the inner ECN onto the outer carrier on encap and
|
||||
// decryptToTun folds outer CE into the inner header on decap. Toggle
|
||||
// via tunnels.ecn (default true).
|
||||
ecnEnabled atomic.Bool
|
||||
pinThreads bool
|
||||
relayManager *relayManager
|
||||
|
||||
tryPromoteEvery atomic.Uint32
|
||||
@@ -111,8 +107,8 @@ type Interface struct {
|
||||
sendRecvErrorConfig recvErrorConfig
|
||||
acceptRecvErrorConfig recvErrorConfig
|
||||
|
||||
// rebindCount is used to decide if an active tunnel should trigger a punch notification through a lighthouse
|
||||
rebindCount int8
|
||||
// Bumped on every udp rebind, tunnels compare it to decide they need a punch from the far side
|
||||
rebindEpoch atomic.Uint32
|
||||
version string
|
||||
|
||||
conntrackCacheTimeout time.Duration
|
||||
@@ -120,10 +116,12 @@ type Interface struct {
|
||||
ctx context.Context
|
||||
writers []udp.Conn
|
||||
queues []tio.Queue
|
||||
// batchers is one per tun queue, wrapping queues[i].
|
||||
// decryptToTun sends plaintext into the batch.RxBatcher;
|
||||
// listenOut calls its Flush at the end of each UDP recvmmsg batch.
|
||||
batchers []batch.RxBatcher
|
||||
// batchers is one per tun queue, wrapping queues[i]. readOutsidePackets
|
||||
// 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.
|
||||
@@ -135,18 +133,13 @@ type Interface struct {
|
||||
metricHandshakes metrics.Histogram
|
||||
messageMetrics *MessageMetrics
|
||||
cachedPacketMetrics *cachedPacketMetrics
|
||||
metricTxDropped metrics.Counter
|
||||
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
type EncWriter interface {
|
||||
SendVia(via *HostInfo,
|
||||
relay *Relay,
|
||||
ad,
|
||||
nb,
|
||||
out []byte,
|
||||
nocopy bool,
|
||||
)
|
||||
SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int)
|
||||
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)
|
||||
Handshake(vpnAddr netip.Addr)
|
||||
@@ -205,6 +198,10 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
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()
|
||||
ifce := &Interface{
|
||||
ctx: ctx,
|
||||
@@ -222,7 +219,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
routines: c.routines,
|
||||
version: c.version,
|
||||
writers: make([]udp.Conn, c.routines),
|
||||
batchers: make([]batch.RxBatcher, c.routines),
|
||||
batchers: make([]*batch.MultiCoalescer, c.routines),
|
||||
myVpnNetworks: cs.myVpnNetworks,
|
||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||
myVpnAddrs: cs.myVpnAddrs,
|
||||
@@ -235,6 +232,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
pinThreads: c.PinThreads,
|
||||
|
||||
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
||||
metricTxDropped: metrics.GetOrRegisterCounter("udp.tx.dropped", nil),
|
||||
messageMetrics: c.MessageMetrics,
|
||||
cachedPacketMetrics: &cachedPacketMetrics{
|
||||
sent: metrics.GetOrRegisterCounter("hostinfo.cached_packets.sent", nil),
|
||||
@@ -273,6 +271,9 @@ func (f *Interface) activate() error {
|
||||
"build", f.version,
|
||||
"udpAddr", addr,
|
||||
"boringcrypto", boringEnabled(),
|
||||
"fips140Version", fips140.Version(),
|
||||
"fips140Enabled", fips140.Enabled(),
|
||||
"fips140Enforced", fips140.Enforced(),
|
||||
)
|
||||
|
||||
if f.routines > 1 && !f.outside.SupportsMultipleReaders() {
|
||||
@@ -288,6 +289,13 @@ func (f *Interface) activate() error {
|
||||
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)
|
||||
@@ -297,18 +305,7 @@ func (f *Interface) activate() error {
|
||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||
|
||||
for i := range f.queues {
|
||||
caps := tio.QueueCapabilities(f.queues[i])
|
||||
if caps.TSO || caps.USO {
|
||||
// Multi-lane: TCP gets coalesced when TSO is on, UDP when USO
|
||||
// is on, everything else (and either lane disabled) falls
|
||||
// through to passthrough so non-IP / non-TCP-UDP traffic still
|
||||
// reaches the TUN.
|
||||
arena := batch.NewArena(batch.DefaultMultiArenaCap)
|
||||
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l, arena, caps.TSO, caps.USO)
|
||||
} else {
|
||||
arena := batch.NewArena(batch.DefaultPassthroughArenaCap)
|
||||
f.batchers[i] = batch.NewPassthrough(f.queues[i], arena.Reserve, arena.Reset)
|
||||
}
|
||||
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l)
|
||||
}
|
||||
|
||||
// On error the caller owns the cleanup, Control.Start cancels the service context
|
||||
@@ -356,6 +353,31 @@ 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) {
|
||||
var li udp.Conn
|
||||
if i > 0 {
|
||||
@@ -364,21 +386,17 @@ func (f *Interface) listenOut(i int) {
|
||||
li = f.outside
|
||||
}
|
||||
|
||||
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||
lhh := f.lightHouse.NewRequestHandler()
|
||||
h := &header.H{}
|
||||
fwPacket := &firewall.Packet{}
|
||||
nb := make([]byte, 12, 12)
|
||||
rxc := newRxContext(f, i)
|
||||
|
||||
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) {
|
||||
plaintext := f.batchers[i].Reserve(len(payload))
|
||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(), meta)
|
||||
listener := func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, payload, rxc)
|
||||
}
|
||||
|
||||
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)
|
||||
@@ -394,32 +412,36 @@ func (f *Interface) listenOut(i int) {
|
||||
f.l.Debug("underlay reader is done", "reader", i)
|
||||
}
|
||||
|
||||
func (f *Interface) pinThisThread(i int) {
|
||||
var cpu int
|
||||
if n := len(f.cpuAffinity); n > 0 {
|
||||
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already
|
||||
// validated the entries against the allowed CPU set.
|
||||
cpu = f.cpuAffinity[i%n]
|
||||
} else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 {
|
||||
// Default: spread queues across the CPUs we're actually allowed to
|
||||
// run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so
|
||||
// i % NumCPU would pick unrunnable IDs and every pin would fail.
|
||||
cpu = allowed[i%len(allowed)]
|
||||
} else {
|
||||
cpu = i % runtime.NumCPU()
|
||||
}
|
||||
if err := util.PinThreadToCPU(cpu); err != nil {
|
||||
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
|
||||
}
|
||||
}
|
||||
|
||||
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 {
|
||||
var cpu int
|
||||
if n := len(f.cpuAffinity); n > 0 {
|
||||
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already
|
||||
// validated the entries against the allowed CPU set.
|
||||
cpu = f.cpuAffinity[i%n]
|
||||
} else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 {
|
||||
// Default: spread queues across the CPUs we're actually allowed to
|
||||
// run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so
|
||||
// i % NumCPU would pick unrunnable IDs and every pin would fail.
|
||||
cpu = allowed[i%len(allowed)]
|
||||
} else {
|
||||
cpu = i % runtime.NumCPU()
|
||||
}
|
||||
if err := util.PinThreadToCPU(cpu); err != nil {
|
||||
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
|
||||
}
|
||||
f.pinThisThread(i)
|
||||
}
|
||||
|
||||
rejectBuf := make([]byte, mtu)
|
||||
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
|
||||
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
|
||||
fwPacket := &firewall.Packet{}
|
||||
fwPacket := &firewall.ParsedPacket{}
|
||||
nb := make([]byte, 12, 12)
|
||||
|
||||
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||
@@ -441,26 +463,35 @@ func (f *Interface) listenIn(queue tio.Queue, i int) {
|
||||
// accumulated so the first packets of a deep read drain
|
||||
// hit the wire while the rest are still being encrypted.
|
||||
if sb.Len() >= batch.SendBatchCap {
|
||||
if err := sb.Flush(); err != nil {
|
||||
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
||||
}
|
||||
f.flushSendBatch(sb, i)
|
||||
}
|
||||
}
|
||||
if err := sb.Flush(); err != nil {
|
||||
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
||||
}
|
||||
f.flushSendBatch(sb, 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) {
|
||||
c.RegisterReloadCallback(f.reloadFirewall)
|
||||
c.RegisterReloadCallback(f.reloadSendRecvError)
|
||||
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
||||
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
||||
c.RegisterReloadCallback(f.reloadMisc)
|
||||
c.RegisterReloadCallback(f.reloadEcn)
|
||||
|
||||
for _, udpConn := range f.writers {
|
||||
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
||||
@@ -593,23 +624,6 @@ func (f *Interface) reloadMisc(c *config.C) {
|
||||
}
|
||||
}
|
||||
|
||||
// reloadEcn syncs Interface.ecnEnabled with the tunnels.ecn config knob.
|
||||
// Default is enabled (RFC 6040 normal mode); set false on the rare path
|
||||
// where an underlay middlebox rewrites or drops ECN bits unpredictably.
|
||||
func (f *Interface) reloadEcn(c *config.C) {
|
||||
initial := c.InitialLoad()
|
||||
if initial || c.HasChanged("tunnels.ecn") {
|
||||
v := c.GetBool("tunnels.ecn", true)
|
||||
changed := f.ecnEnabled.Swap(v) != v
|
||||
if !initial {
|
||||
f.l.Info("tunnels.ecn changed", "enabled", v)
|
||||
if changed {
|
||||
f.l.Warn("tunnels.ecn datapath toggled, but route-level ECN negotiation (RTAX_FEATURE_ECN) retains its previous state until nebula is restarted", "enabled", v)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
||||
ticker := time.NewTicker(i)
|
||||
defer ticker.Stop()
|
||||
|
||||
@@ -0,0 +1,146 @@
|
||||
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)
|
||||
}
|
||||
@@ -0,0 +1,242 @@
|
||||
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))
|
||||
}
|
||||
}
|
||||
+41
-9
@@ -2,11 +2,16 @@ package iputil
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
|
||||
"golang.org/x/net/ipv4"
|
||||
"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 (
|
||||
// MaxIPv4RejectPacketSize is the largest IPv4 reject packet:
|
||||
// - 20 byte ipv4 header
|
||||
@@ -22,6 +27,13 @@ const (
|
||||
maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000
|
||||
|
||||
MaxRejectPacketSize = maxIPv6RejectPacketSize
|
||||
|
||||
IPProtocolICMP = 1
|
||||
IPProtocolICMPv6 = 58
|
||||
IPProtocolTCP = 6
|
||||
IPProtocolUDP = 17
|
||||
ICMPv6TypeEchoRequest = 128
|
||||
ICMPv6TypeEchoReply = 129
|
||||
)
|
||||
|
||||
func CreateRejectPacket(packet []byte, out []byte) []byte {
|
||||
@@ -199,8 +211,8 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte {
|
||||
}
|
||||
|
||||
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
|
||||
proto, offset, isFragment := ipv6FindUpperProtocol(packet)
|
||||
if isFragment {
|
||||
proto, offset, isFragment, _, err := IPv6FindUpperProtocol(packet)
|
||||
if err != nil || isFragment {
|
||||
return nil
|
||||
}
|
||||
switch proto {
|
||||
@@ -333,40 +345,60 @@ func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte {
|
||||
return out
|
||||
}
|
||||
|
||||
func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
|
||||
// IPv6FindUpperProtocol walks the ipv6 extension header chain and returns the upper layer protocol, the
|
||||
// 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]
|
||||
offset = ipv6.HeaderLen
|
||||
|
||||
for {
|
||||
for range maxIPv6ExtHeaders {
|
||||
switch nextHeader {
|
||||
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
|
||||
if len(packet) < offset+2 {
|
||||
return nextHeader, offset, isFragment
|
||||
return nextHeader, offset, isFragment, anyFragment, ErrIPv6CouldNotFindPayload
|
||||
}
|
||||
nextHeader = packet[offset]
|
||||
offset += (int(packet[offset+1]) + 1) << 3
|
||||
|
||||
case 44: // Fragment
|
||||
if len(packet) < offset+8 {
|
||||
return nextHeader, offset, isFragment
|
||||
return nextHeader, offset, isFragment, anyFragment, ErrIPv6CouldNotFindPayload
|
||||
}
|
||||
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 {
|
||||
isFragment = true
|
||||
return packet[offset], offset, true, anyFragment, nil
|
||||
}
|
||||
nextHeader = packet[offset]
|
||||
offset += 8
|
||||
|
||||
case 51: // AH
|
||||
if len(packet) < offset+2 {
|
||||
return nextHeader, offset, isFragment
|
||||
return nextHeader, offset, isFragment, anyFragment, ErrIPv6CouldNotFindPayload
|
||||
}
|
||||
nextHeader = packet[offset]
|
||||
offset += (int(packet[offset+1]) + 2) << 2
|
||||
|
||||
default:
|
||||
return nextHeader, offset, isFragment
|
||||
// A prior extension header can declare a length that advances offset past the packet. The terminal
|
||||
// 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 {
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/net/ipv4"
|
||||
"golang.org/x/net/ipv6"
|
||||
)
|
||||
@@ -515,3 +516,63 @@ func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) {
|
||||
result := CreateICMPEchoResponse(packet, out)
|
||||
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)
|
||||
})
|
||||
}
|
||||
|
||||
+24
-2
@@ -34,7 +34,13 @@ type LightHouse struct {
|
||||
|
||||
myVpnNetworks []netip.Prefix
|
||||
myVpnNetworksTable *bart.Lite
|
||||
punchy *Punchy
|
||||
// myVpnAddrsTable contains our overlay host addrs, as opposed to the overlay networks
|
||||
myVpnAddrsTable *bart.Lite
|
||||
punchy *Punchy
|
||||
|
||||
// localAddrsFn enumerates the underlay addresses we advertise. It is a field so tests can supply simulated
|
||||
// addresses rather than whatever this machine's NICs happen to be. Set it before Start.
|
||||
localAddrsFn func(*LocalAllowList) []netip.Addr
|
||||
|
||||
// Local cache of answers from light houses
|
||||
// map of vpn addr to answers
|
||||
@@ -100,6 +106,7 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
||||
amLighthouse: amLighthouse,
|
||||
myVpnNetworks: cs.myVpnNetworks,
|
||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||
myVpnAddrsTable: cs.myVpnAddrsTable,
|
||||
addrMap: make(map[netip.Addr]*RemoteList),
|
||||
nebulaPort: nebulaPort,
|
||||
punchy: p,
|
||||
@@ -107,6 +114,10 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
||||
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
||||
l: l,
|
||||
}
|
||||
h.localAddrsFn = func(al *LocalAllowList) []netip.Addr {
|
||||
return localAddrs(h.l, al)
|
||||
}
|
||||
|
||||
lighthouses := make([]netip.Addr, 0)
|
||||
h.lighthouses.Store(&lighthouses)
|
||||
staticList := make(map[netip.Addr]struct{})
|
||||
@@ -918,7 +929,7 @@ func (lh *LightHouse) SendUpdate() {
|
||||
}
|
||||
|
||||
lal := lh.GetLocalAllowList()
|
||||
for _, e := range localAddrs(lh.l, lal) {
|
||||
for _, e := range lh.localAddrsFn(lal) {
|
||||
if lh.myVpnNetworksTable.Contains(e) {
|
||||
continue
|
||||
}
|
||||
@@ -1150,6 +1161,17 @@ func (lhh *LightHouseHandler) handleHostQuery(n *NebulaMeta, fromVpnAddrs []neti
|
||||
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) {
|
||||
n = lhh.resetMeta()
|
||||
n.Type = NebulaMeta_HostQueryReply
|
||||
|
||||
+81
-55
@@ -27,15 +27,27 @@ func TestOldIPv4Only(t *testing.T) {
|
||||
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) {
|
||||
l := test.NewLogger()
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
cs := testCertState(myVpnNet)
|
||||
lh1 := "10.128.0.2"
|
||||
|
||||
c := config.NewC(l)
|
||||
@@ -55,12 +67,7 @@ func Test_lhStaticMapping(t *testing.T) {
|
||||
func TestReloadLighthouseInterval(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
cs := testCertState(myVpnNet)
|
||||
lh1 := "10.128.0.2"
|
||||
|
||||
c := config.NewC(l)
|
||||
@@ -90,12 +97,7 @@ func TestReloadLighthouseInterval(t *testing.T) {
|
||||
func BenchmarkLighthouseHandleRequest(b *testing.B) {
|
||||
l := test.NewLogger()
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/0")
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
cs := testCertState(myVpnNet)
|
||||
|
||||
c := config.NewC(l)
|
||||
lh, err := NewLightHouseFromConfig(b.Context(), l, c, cs, nil, nil)
|
||||
@@ -195,12 +197,7 @@ func TestLighthouse_Memory(t *testing.T) {
|
||||
c.Settings["listen"] = map[string]any{"port": 4242}
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
cs := testCertState(myVpnNet)
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||
lh.ifce = &mockEncWriter{}
|
||||
require.NoError(t, err)
|
||||
@@ -280,12 +277,7 @@ func TestLighthouse_reload(t *testing.T) {
|
||||
c.Settings["listen"] = map[string]any{"port": 4242}
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
cs := testCertState(myVpnNet)
|
||||
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||
require.NoError(t, err)
|
||||
@@ -315,12 +307,7 @@ func TestLighthouse_reloadStaticHostMap(t *testing.T) {
|
||||
}
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
cs := testCertState(myVpnNet)
|
||||
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||
require.NoError(t, err)
|
||||
@@ -429,7 +416,9 @@ func TestLighthouse_reloadStaticHostMap(t *testing.T) {
|
||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
||||
}
|
||||
|
||||
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
||||
// sendLHHostRequest delivers a HostQuery to lhh and hands back the writer that
|
||||
// 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{
|
||||
Type: NebulaMeta_HostQuery,
|
||||
Details: &NebulaMetaDetails{},
|
||||
@@ -447,12 +436,59 @@ func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, l
|
||||
panic(err)
|
||||
}
|
||||
|
||||
filter := NebulaMeta_HostQueryReply
|
||||
w := &testEncWriter{
|
||||
metaFilter: &filter,
|
||||
}
|
||||
w := &testEncWriter{metaFilter: filter}
|
||||
lhh.HandleRequest(fromAddr, []netip.Addr{myVpnIp}, b, w)
|
||||
return w.lastReply
|
||||
return w
|
||||
}
|
||||
|
||||
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
||||
filter := NebulaMeta_HostQueryReply
|
||||
return sendLHHostRequest(fromAddr, myVpnIp, queryVpnIp, lhh, &filter).lastReply
|
||||
}
|
||||
|
||||
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"},
|
||||
}
|
||||
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, testCertState(myVpnNet), nil, nil)
|
||||
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) {
|
||||
@@ -498,7 +534,7 @@ type testEncWriter struct {
|
||||
protocolVersion cert.Version
|
||||
}
|
||||
|
||||
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool) {
|
||||
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
||||
}
|
||||
func (tw *testEncWriter) Handshake(vpnIp netip.Addr) {
|
||||
}
|
||||
@@ -642,12 +678,7 @@ func TestLighthouse_Dont_Delete_Static_Hosts(t *testing.T) {
|
||||
}
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
cs := testCertState(myVpnNet)
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||
require.NoError(t, err)
|
||||
lh.ifce = &mockEncWriter{}
|
||||
@@ -708,12 +739,7 @@ func TestLighthouse_DeletesWork(t *testing.T) {
|
||||
}
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
cs := testCertState(myVpnNet)
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||
require.NoError(t, err)
|
||||
lh.ifce = &mockEncWriter{}
|
||||
|
||||
@@ -6,12 +6,15 @@ import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"runtime/debug"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/cpupick"
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
"github.com/slackhq/nebula/overlay"
|
||||
"github.com/slackhq/nebula/sshd"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
@@ -21,6 +24,12 @@ import (
|
||||
|
||||
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) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
// Automatically cancel the context if Main returns an error, to signal all created goroutines to quit.
|
||||
@@ -85,9 +94,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
if routines < 1 {
|
||||
routines = 1
|
||||
}
|
||||
if routines > 1 {
|
||||
l.Info("Using multiple routines", "routines", routines)
|
||||
}
|
||||
} else {
|
||||
// deprecated and undocumented
|
||||
tunQueues := c.GetInt("tun.routines", 1)
|
||||
@@ -97,6 +103,12 @@ 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)
|
||||
}
|
||||
}
|
||||
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
|
||||
// Intentionally not documented yet while we do more testing and determine
|
||||
@@ -164,8 +176,21 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
}
|
||||
|
||||
for i := 0; i < routines; i++ {
|
||||
l.Info("listening", "addr", netip.AddrPortFrom(listenHost, uint16(port)))
|
||||
udpServer, err := udp.NewListener(l, listenHost, port, routines > 1, c.GetInt("listen.batch", 64))
|
||||
listen := netip.AddrPortFrom(listenHost, uint16(port))
|
||||
l.Info("listening", "addr", listen)
|
||||
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 {
|
||||
return nil, util.NewContextualError("Failed to open udp listener", m{"queue": i}, err)
|
||||
}
|
||||
@@ -216,8 +241,33 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
|
||||
pinThreads := c.GetBool("tun.pin_threads", true)
|
||||
cpuAffinity := parseCpuAffinity(c, l, routines)
|
||||
if pinThreads && len(cpuAffinity) == 0 && !configTest {
|
||||
cpuAffinity = defaultCPUAffinityAvoidingIRQs(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{
|
||||
@@ -260,7 +310,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
ifce.reloadDisconnectInvalid(c)
|
||||
ifce.reloadSendRecvError(c)
|
||||
ifce.reloadAcceptRecvError(c)
|
||||
ifce.reloadEcn(c)
|
||||
|
||||
handshakeManager.f = ifce
|
||||
go handshakeManager.Run(ctx)
|
||||
@@ -281,6 +330,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
|
||||
attachCommands(l, c, ssh, ifce)
|
||||
|
||||
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
|
||||
|
||||
return &Control{
|
||||
state: StateReady,
|
||||
f: ifce,
|
||||
@@ -291,6 +342,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
statsStart: stats.Start,
|
||||
dnsStart: ds.Start,
|
||||
lighthouseStart: lightHouse.StartUpdateWorker,
|
||||
networkChangeStart: networkChanges.Start,
|
||||
connectionManagerStart: connManager.Start,
|
||||
}, nil
|
||||
}
|
||||
@@ -359,57 +411,6 @@ func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
|
||||
return cpus
|
||||
}
|
||||
|
||||
// defaultCPUAffinityAvoidingIRQs picks the default pin set for the tun
|
||||
// readers when tun.cpu_affinity is unset: allowed CPUs that do NOT service
|
||||
// any physical NIC's interrupts. The stock allowed[i] spread pins the
|
||||
// encrypt threads onto exactly the cores most drivers affine their first RX
|
||||
// queue IRQs to, so whenever a flow's RSS queue fires on a core hosting a
|
||||
// tun reader, NAPI and encrypt fight for the core and per-flow throughput
|
||||
// drops (measured: REV 8.4 vs 10.2 Gbps on the same hardware, 2026-07-14).
|
||||
//
|
||||
// Returns nil — keeping the old allowed[i] fallback in listenIn — when IRQ
|
||||
// info is unavailable or when there aren't enough IRQ-free CPUs to give
|
||||
// every routine its own core: silently doubling readers up on fewer cores
|
||||
// is worse than the occasional IRQ collision. NICs whose vectors blanket
|
||||
// every CPU (e.g. mlx5 defaults to one queue per core) make avoidance
|
||||
// impossible; narrowing the NIC's spread (ethtool -X <dev> equal N, or
|
||||
// /proc/irq/*/smp_affinity) or setting tun.cpu_affinity explicitly makes it
|
||||
// effective.
|
||||
func defaultCPUAffinityAvoidingIRQs(l *slog.Logger, routines int) []int {
|
||||
irq, err := util.NICIRQCPUs()
|
||||
if err != nil || len(irq) == 0 {
|
||||
return nil
|
||||
}
|
||||
allowed, err := util.AllowedCPUs()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
cpus := chooseIRQFreeCPUs(allowed, irq, routines)
|
||||
if cpus == nil {
|
||||
l.Info("not enough CPUs are free of NIC IRQs to give every tun reader its own; using the default spread",
|
||||
"routines", routines, "allowed", len(allowed), "irqCPUs", len(irq))
|
||||
return nil
|
||||
}
|
||||
l.Info("pinning tun readers to CPUs clear of NIC IRQs", "cpus", cpus)
|
||||
return cpus
|
||||
}
|
||||
|
||||
// chooseIRQFreeCPUs returns the first `routines` allowed CPUs not present in
|
||||
// irq, or nil if fewer than `routines` qualify.
|
||||
func chooseIRQFreeCPUs(allowed []int, irq map[int]bool, routines int) []int {
|
||||
free := make([]int, 0, routines)
|
||||
for _, cpu := range allowed {
|
||||
if irq[cpu] {
|
||||
continue
|
||||
}
|
||||
free = append(free, cpu)
|
||||
if len(free) == routines {
|
||||
return free
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func moduleVersion() string {
|
||||
info, ok := debug.ReadBuildInfo()
|
||||
if !ok {
|
||||
|
||||
@@ -9,26 +9,6 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestChooseIRQFreeCPUs(t *testing.T) {
|
||||
irq := map[int]bool{0: true, 1: true, 2: true, 3: true}
|
||||
|
||||
// Plenty of IRQ-free CPUs: take the first `routines` of them in order.
|
||||
assert.Equal(t, []int{4, 5}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 2))
|
||||
|
||||
// Exactly enough.
|
||||
assert.Equal(t, []int{4, 5, 6}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 3))
|
||||
|
||||
// Not enough IRQ-free CPUs: nil, caller keeps the old default rather
|
||||
// than doubling readers up on shared cores.
|
||||
assert.Nil(t, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4}, irq, 2))
|
||||
|
||||
// No IRQ info at all behaves like a plain prefix of allowed.
|
||||
assert.Equal(t, []int{0, 1}, chooseIRQFreeCPUs([]int{0, 1, 2}, map[int]bool{}, 2))
|
||||
|
||||
// Non-contiguous allowed set (cgroup cpuset) with holes.
|
||||
assert.Equal(t, []int{9, 12}, chooseIRQFreeCPUs([]int{1, 3, 9, 12}, map[int]bool{1: true, 3: true}, 2))
|
||||
}
|
||||
|
||||
func TestParseCpuAffinity(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
|
||||
|
||||
+13
-4
@@ -14,7 +14,8 @@ type MessageMetrics struct {
|
||||
rxUnknown 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) {
|
||||
@@ -41,6 +42,13 @@ 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 {
|
||||
gen := func(t string) [][]metrics.Counter {
|
||||
return [][]metrics.Counter{
|
||||
@@ -61,9 +69,10 @@ func newMessageMetrics() *MessageMetrics {
|
||||
rx: gen("rx"),
|
||||
tx: gen("tx"),
|
||||
|
||||
rxUnknown: metrics.GetOrRegisterCounter("messages.rx.other", nil),
|
||||
txUnknown: metrics.GetOrRegisterCounter("messages.tx.other", nil),
|
||||
rxInvalid: metrics.GetOrRegisterCounter("messages.rx.invalid", nil),
|
||||
rxUnknown: metrics.GetOrRegisterCounter("messages.rx.other", nil),
|
||||
txUnknown: metrics.GetOrRegisterCounter("messages.tx.other", nil),
|
||||
rxInvalid: metrics.GetOrRegisterCounter("messages.rx.invalid", nil),
|
||||
txExhausted: metrics.GetOrRegisterCounter("messages.tx.exhausted", nil),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -25,6 +25,9 @@ func (s *CipherStateAESGCM) EncryptDanger(out, ad, plaintext []byte, n uint64, n
|
||||
if s == nil {
|
||||
return nil, errors.New("no cipher state available to encrypt")
|
||||
}
|
||||
if n >= RejectAfterMessages {
|
||||
return nil, ErrMessageCounterExhausted
|
||||
}
|
||||
nb[0] = 0
|
||||
nb[1] = 0
|
||||
nb[2] = 0
|
||||
|
||||
+4
-65
@@ -4,77 +4,16 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"encoding/binary"
|
||||
|
||||
// unsafe needed for go:linkname
|
||||
_ "unsafe"
|
||||
"crypto/boring"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
)
|
||||
|
||||
var CipherAESGCM noise.CipherFunc = CipherAESGCMFIPS140
|
||||
|
||||
// EncryptLockNeeded indicates if calls to Encrypt need a lock
|
||||
// This is true for boringcrypto because the Seal function verifies that the
|
||||
// nonce is strictly increasing.
|
||||
const EncryptLockNeeded = true
|
||||
|
||||
// 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)
|
||||
}
|
||||
var boringEnabled = boring.Enabled()
|
||||
|
||||
@@ -4,8 +4,6 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/boring"
|
||||
"encoding/hex"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -14,33 +12,3 @@ import (
|
||||
func TestEncryptLockNeeded(t *testing.T) {
|
||||
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,6 +24,9 @@ func (s *CipherStateChaChaPoly) EncryptDanger(out, ad, plaintext []byte, n uint6
|
||||
if s == nil {
|
||||
return nil, errors.New("no cipher state available to encrypt")
|
||||
}
|
||||
if n >= RejectAfterMessages {
|
||||
return nil, ErrMessageCounterExhausted
|
||||
}
|
||||
nb[0] = 0
|
||||
nb[1] = 0
|
||||
nb[2] = 0
|
||||
|
||||
@@ -1,11 +1,22 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
|
||||
"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.
|
||||
// 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.
|
||||
@@ -29,8 +40,11 @@ type CipherState interface {
|
||||
// NewCipherState wraps the post-handshake noise.CipherState in the per-cipher type that matches cipherFunc.
|
||||
// cipherFunc must be the same cipher used to build the noise CipherSuite that produced s.
|
||||
func NewCipherState(s *noise.CipherState, cipherFunc noise.CipherFunc) CipherState {
|
||||
if cs, ok := s.Cipher().(CipherState); ok {
|
||||
return cs
|
||||
}
|
||||
switch cipherFunc.CipherName() {
|
||||
case CipherAESGCM.CipherName():
|
||||
case noise.CipherAESGCM.CipherName():
|
||||
return NewCipherStateAESGCM(s)
|
||||
case noise.CipherChaChaPoly.CipherName():
|
||||
return NewCipherStateChaChaPoly(s)
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/fips140"
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
@@ -10,24 +12,30 @@ import (
|
||||
|
||||
func TestCipherStateAESGCMRoundtrip(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, CipherAESGCM)
|
||||
roundtrip(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec))
|
||||
roundtrip(t, NewCipherState(enc, CipherAESGCM), NewCipherState(dec, CipherAESGCM))
|
||||
}
|
||||
|
||||
func TestCipherStateChaChaPolyRoundtrip(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
roundtrip(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec))
|
||||
roundtrip(t, NewCipherState(enc, noise.CipherChaChaPoly), NewCipherState(dec, noise.CipherChaChaPoly))
|
||||
}
|
||||
|
||||
func TestNewCipherStateDispatch(t *testing.T) {
|
||||
encA, _ := buildCipherStates(t, CipherAESGCM)
|
||||
encC, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
|
||||
assert.IsType(t, &CipherStateAESGCM{}, NewCipherState(encA, CipherAESGCM))
|
||||
if !boringEnabled && !fips140.Enabled() {
|
||||
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))
|
||||
}
|
||||
|
||||
func TestNewCipherStateUnsupportedPanics(t *testing.T) {
|
||||
enc, _ := buildCipherStates(t, CipherAESGCM)
|
||||
enc, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
assert.Panics(t, func() {
|
||||
NewCipherState(enc, fakeCipher{})
|
||||
})
|
||||
@@ -89,6 +97,24 @@ func roundtrip(t *testing.T, enc, dec CipherState) {
|
||||
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) {
|
||||
enc, _ := buildCipherStatesB(b, CipherAESGCM)
|
||||
benchEncryptCipherState(b, NewCipherState(enc, CipherAESGCM))
|
||||
@@ -164,3 +190,48 @@ func TestCipherStateNilSafety(t *testing.T) {
|
||||
assert.Empty(t, out)
|
||||
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")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,197 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
//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")
|
||||
}
|
||||
}
|
||||
+15
-4
@@ -1,14 +1,25 @@
|
||||
//go:build !boringcrypto
|
||||
// +build !boringcrypto
|
||||
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/fips140"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
)
|
||||
|
||||
// EncryptLockNeeded indicates if calls to Encrypt need a lock
|
||||
const EncryptLockNeeded = false
|
||||
var EncryptLockNeeded = fips140.Enabled()
|
||||
|
||||
// CipherAESGCM is the standard noise.CipherAESGCM when boringcrypto is not enabled
|
||||
var CipherAESGCM noise.CipherFunc = noise.CipherAESGCM
|
||||
var CipherAESGCM noise.CipherFunc = initAESGCM()
|
||||
|
||||
func initAESGCM() noise.CipherFunc {
|
||||
if fips140.Enabled() {
|
||||
return CipherAESGCMFIPS140
|
||||
} else {
|
||||
return noise.CipherAESGCM
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
var boringEnabled = false
|
||||
|
||||
@@ -1,14 +0,0 @@
|
||||
//go:build !boringcrypto
|
||||
// +build !boringcrypto
|
||||
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestEncryptLockNeeded(t *testing.T) {
|
||||
assert.False(t, EncryptLockNeeded)
|
||||
}
|
||||
+114
-242
@@ -8,12 +8,12 @@ import (
|
||||
"net/netip"
|
||||
"time"
|
||||
|
||||
"github.com/google/gopacket/layers"
|
||||
"golang.org/x/net/ipv6"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
"github.com/slackhq/nebula/overlay/batch"
|
||||
"golang.org/x/net/ipv4"
|
||||
)
|
||||
|
||||
@@ -23,7 +23,11 @@ const (
|
||||
|
||||
var ErrOutOfWindow = errors.New("out of window packet")
|
||||
|
||||
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||
// readOutsidePackets processes one received underlay packet.
|
||||
// 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)
|
||||
if err != nil {
|
||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
||||
@@ -91,7 +95,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
if isMessageRelay {
|
||||
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
|
||||
} else {
|
||||
hostinfo = f.hostMap.QueryIndex(h.RemoteIndex)
|
||||
hostinfo = f.hostMap.QueryIndexCached(h.RemoteIndex, rxc.hostmapCache)
|
||||
}
|
||||
|
||||
// At this point we should have a valid existing tunnel, verify and send
|
||||
@@ -103,26 +107,32 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
return
|
||||
}
|
||||
|
||||
if len(packet) < header.Len+hostinfo.ConnectionState.dKey.Overhead() {
|
||||
f.messageMetrics.RxInvalid(1)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("packet too small", "from", via, "length", len(packet))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// All remaining packets are encrypted
|
||||
ci := hostinfo.ConnectionState
|
||||
if !ci.window.Check(f.l, h.MessageCounter) {
|
||||
return
|
||||
}
|
||||
|
||||
// Relay packets are special
|
||||
if isMessageRelay {
|
||||
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache, meta)
|
||||
// Relay packets are special, this branch should always early-return
|
||||
err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, rxc.nb)
|
||||
if err != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h)
|
||||
}
|
||||
return
|
||||
}
|
||||
f.handleOutsideRelayPacket(hostinfo, via, packet, rxc)
|
||||
return
|
||||
}
|
||||
|
||||
out, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||
out, err := hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, packet, rxc.nb)
|
||||
if err != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Failed to decrypt packet",
|
||||
"error", err,
|
||||
"from", via,
|
||||
"header", h,
|
||||
)
|
||||
hostinfo.logger(f.l).Debug("Failed to decrypt packet", "error", err, "from", via, "header", h)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -135,7 +145,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
case header.Message:
|
||||
switch h.Subtype {
|
||||
case header.MessageNone:
|
||||
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache, meta)
|
||||
f.handleOutsideMessagePacket(hostinfo, h.MessageCounter, out, rxc)
|
||||
default:
|
||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
||||
return
|
||||
@@ -143,15 +153,23 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
|
||||
case header.LightHouse:
|
||||
//TODO: assert via is not relayed
|
||||
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
||||
rxc.lhh.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
||||
|
||||
case header.Test:
|
||||
switch h.Subtype {
|
||||
case header.TestReply:
|
||||
// No-op, useful for the Roaming and connectionManager side-effects above
|
||||
case header.TestRequest:
|
||||
//recycle the input packet ciphertext as our output buffer
|
||||
f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, packet)
|
||||
const maxCipherOverhead = 16 //todo we use this too often, needs a real importable const
|
||||
const maxOverhead = header.Len + header.Len + maxCipherOverhead + maxCipherOverhead
|
||||
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:
|
||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
|
||||
return
|
||||
@@ -169,28 +187,10 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||
// The entire body is sent as AD, not encrypted.
|
||||
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
||||
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
|
||||
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
|
||||
// which will gracefully fail in the DecryptDanger call.
|
||||
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
||||
var err error
|
||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
// Advance the replay window now that the frame is authenticated
|
||||
if !hostinfo.ConnectionState.window.Update(f.l, h.MessageCounter) {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("dropping out of window relay packet", "header", h)
|
||||
}
|
||||
return
|
||||
}
|
||||
// Successfully validated the thing. Get rid of the Relay header.
|
||||
signedPayload = signedPayload[header.Len:]
|
||||
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, packet []byte, rxc *rxContext) {
|
||||
h := rxc.h
|
||||
// Successfully validated the thing. Get rid of the Relay header and the AEAD tag
|
||||
signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||
// Pull the Roaming parts up here, and return in all call paths.
|
||||
f.handleHostRoaming(hostinfo, via)
|
||||
// Track usage of both the HostInfo and the Relay for the received & authenticated packet
|
||||
@@ -201,9 +201,7 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
||||
if !ok {
|
||||
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
|
||||
// its internal mapping. This should never happen.
|
||||
hostinfo.logger(f.l).Error("HostInfo missing remote relay index",
|
||||
"relayRemoteIndex", h.RemoteIndex,
|
||||
)
|
||||
hostinfo.logger(f.l).Error("HostInfo missing remote relay index", "relayRemoteIndex", h.RemoteIndex)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -214,11 +212,10 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
||||
via = ViaSender{
|
||||
UdpAddr: via.UdpAddr,
|
||||
relayHI: hostinfo,
|
||||
remoteIdx: relay.RemoteIndex,
|
||||
relay: relay,
|
||||
IsRelayed: true,
|
||||
}
|
||||
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache, meta)
|
||||
f.readOutsidePackets(via, signedPayload, rxc)
|
||||
case ForwardingType:
|
||||
// Find the target HostInfo relay object
|
||||
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
||||
@@ -235,9 +232,11 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
||||
if targetRelay.State == Established {
|
||||
switch targetRelay.Type {
|
||||
case ForwardingType:
|
||||
// Forward this packet through the relay tunnel
|
||||
// Find the target HostInfo //todo it would potentially be nice to batch these
|
||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
||||
// 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
|
||||
fwdBuf := packet[:0]
|
||||
//todo it would potentially be nice to batch these
|
||||
f.SendVia(targetHI, targetRelay, signedPayload, rxc.nb, fwdBuf, true, rxc.q)
|
||||
case TerminalType:
|
||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||
return
|
||||
@@ -314,11 +313,14 @@ var (
|
||||
ErrIPv4InvalidHeaderLength = errors.New("invalid ipv4 header length")
|
||||
ErrIPv4PacketTooShort = errors.New("ipv4 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
|
||||
func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
func newPacket(data []byte, incoming bool, fp *firewall.ParsedPacket) 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 {
|
||||
return ErrPacketTooShort
|
||||
}
|
||||
@@ -333,7 +335,7 @@ func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
return ErrUnknownIPVersion
|
||||
}
|
||||
|
||||
func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
dataLen := len(data)
|
||||
if dataLen < ipv6.HeaderLen {
|
||||
return ErrIPv6PacketTooShort
|
||||
@@ -347,104 +349,64 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
fp.RemoteAddr, _ = netip.AddrFromSlice(data[24:40])
|
||||
}
|
||||
|
||||
protoAt := 6 // NextHeader is at 6 bytes into the ipv6 header
|
||||
offset := ipv6.HeaderLen // Start at the end of the ipv6 header
|
||||
next := 0
|
||||
for {
|
||||
if protoAt >= dataLen {
|
||||
break
|
||||
// Walk the extension header chain to the upper layer protocol. iputil.IPv6FindUpperProtocol is the single
|
||||
// source of truth for which headers are extension headers, so this stays in lockstep with the reject path
|
||||
// and cannot drift into misreading an unknown protocol (SCTP, GRE, etc.) as a forged transport.
|
||||
proto, offset, isFragment, anyFragment, err := iputil.IPv6FindUpperProtocol(data)
|
||||
if err != nil {
|
||||
return ErrIPv6PacketTooShort
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
proto := layers.IPProtocol(data[protoAt])
|
||||
|
||||
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:
|
||||
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
|
||||
switch data[offset] { //icmp type
|
||||
case iputil.ICMPv6TypeEchoRequest, iputil.ICMPv6TypeEchoReply:
|
||||
if dataLen < offset+6 {
|
||||
return ErrIPv6PacketTooShort
|
||||
}
|
||||
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
|
||||
|
||||
fp.RemotePort = binary.BigEndian.Uint16(data[offset+4 : offset+6]) //identifier
|
||||
default:
|
||||
// Normal ipv6 header length processing
|
||||
if dataLen <= offset+1 {
|
||||
break
|
||||
}
|
||||
next = (int(data[offset+1]) + 1) << 3
|
||||
fp.RemotePort = 0
|
||||
}
|
||||
|
||||
if next <= 0 {
|
||||
// Safety check, each ipv6 header has to be at least 8 bytes
|
||||
next = 8
|
||||
case iputil.IPProtocolTCP, iputil.IPProtocolUDP:
|
||||
if dataLen < offset+4 {
|
||||
return ErrIPv6PacketTooShort
|
||||
}
|
||||
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])
|
||||
}
|
||||
|
||||
protoAt = offset
|
||||
offset = offset + next
|
||||
default:
|
||||
// don't set ports for protocols Nebula doesn't inspect
|
||||
fp.RemotePort = 0
|
||||
fp.LocalPort = 0
|
||||
}
|
||||
|
||||
return ErrIPv6CouldNotFindPayload
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
// Do we at least have an ipv4 header worth of data?
|
||||
if len(data) < ipv4.HeaderLen {
|
||||
return ErrIPv4PacketTooShort
|
||||
@@ -461,6 +423,10 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
// Check if this is the second or further fragment of a fragmented packet.
|
||||
flagsfrags := binary.BigEndian.Uint16(data[6:8])
|
||||
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
|
||||
fp.Protocol = data[9]
|
||||
@@ -468,7 +434,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
// Accounting for a variable header length, do we have enough data for our src/dst tuples?
|
||||
minLen := ihl
|
||||
if !fp.Fragment {
|
||||
if fp.Protocol == firewall.ProtoICMP {
|
||||
if fp.Protocol == iputil.IPProtocolICMP {
|
||||
minLen += minFwPacketLen + 2
|
||||
} else {
|
||||
minLen += minFwPacketLen
|
||||
@@ -490,7 +456,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
if fp.Fragment {
|
||||
fp.RemotePort = 0
|
||||
fp.LocalPort = 0
|
||||
} else if fp.Protocol == firewall.ProtoICMP { //note that orientation doesn't matter on ICMP
|
||||
} else if fp.Protocol == iputil.IPProtocolICMP { //note that orientation doesn't matter on ICMP
|
||||
fp.RemotePort = binary.BigEndian.Uint16(data[ihl+4 : ihl+6]) //identifier
|
||||
fp.LocalPort = 0 //code would be uint16(data[ihl+1])
|
||||
} else if incoming {
|
||||
@@ -504,117 +470,23 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []byte, h *header.H, nb []byte) ([]byte, error) {
|
||||
var err error
|
||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], mc, nb)
|
||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounter uint64, out []byte, rxc *rxContext) {
|
||||
err := newPacket(out, true, rxc.fwPacket)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !hostinfo.ConnectionState.window.Update(f.l, mc) {
|
||||
return nil, ErrOutOfWindow
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// 2-bit IP-level ECN codepoints (lower bits of IPv4 ToS / IPv6 TC).
|
||||
const (
|
||||
ecnNotECT = 0x00
|
||||
ecnECT1 = 0x01
|
||||
ecnECT0 = 0x02
|
||||
ecnCE = 0x03
|
||||
)
|
||||
|
||||
// applyOuterECN folds an outer CE mark from the underlay into the inner
|
||||
// IP header per RFC 6040 normal mode. It mutates pkt[1] in place. Other
|
||||
// codepoints are advisory only and leave the inner unchanged.
|
||||
//
|
||||
// Merge cases (outer × inner → action):
|
||||
//
|
||||
// outer != CE : no-op (inner is authoritative)
|
||||
// outer == CE, inner Not-ECT : log; cannot propagate to a non-ECN host
|
||||
// outer == CE, inner ECT/CE : rewrite inner ECN to CE
|
||||
func applyOuterECN(pkt []byte, outerECN byte, hostinfo *HostInfo, l *slog.Logger) {
|
||||
if outerECN&ecnCE != ecnCE || len(pkt) < 2 {
|
||||
return
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
switch pkt[1] & 0x03 {
|
||||
case ecnNotECT:
|
||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
||||
}
|
||||
case ecnCE:
|
||||
// Already CE.
|
||||
default:
|
||||
// Rewriting the ToS byte invalidates the IPv4 header checksum, so
|
||||
// patch it incrementally per RFC 1624 (HC' = ~(~HC + ~m + m')). The
|
||||
// ToS is the low byte of the 16-bit word at pkt[0:2]; the header
|
||||
// checksum lives at pkt[10:12]. A header too short to carry a
|
||||
// checksum can't be fixed up here, so leave it for newPacket to
|
||||
// reject rather than emit a mangled packet.
|
||||
if len(pkt) < ipv4.HeaderLen {
|
||||
return
|
||||
}
|
||||
m := binary.BigEndian.Uint16(pkt[0:2])
|
||||
pkt[1] = (pkt[1] &^ 0x03) | ecnCE
|
||||
mNew := binary.BigEndian.Uint16(pkt[0:2])
|
||||
sum := uint32(^binary.BigEndian.Uint16(pkt[10:12])) + uint32(^m) + uint32(mNew)
|
||||
for sum > 0xffff {
|
||||
sum = (sum >> 16) + (sum & 0xffff)
|
||||
}
|
||||
binary.BigEndian.PutUint16(pkt[10:12], ^uint16(sum))
|
||||
}
|
||||
case 6:
|
||||
switch (pkt[1] >> 4) & 0x03 {
|
||||
case ecnNotECT:
|
||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
||||
}
|
||||
case ecnCE:
|
||||
// Already CE.
|
||||
default:
|
||||
pkt[1] = (pkt[1] &^ 0x30) | (ecnCE << 4)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||
// RFC 6040 normal-mode combine: fold any outer CE mark stamped by the
|
||||
// underlay into the inner header before firewall + TUN write. Other
|
||||
// outer codepoints are advisory only — we keep the inner unchanged.
|
||||
if f.ecnEnabled.Load() {
|
||||
applyOuterECN(out, meta.OuterECN, hostinfo, f.l)
|
||||
}
|
||||
|
||||
err := newPacket(out, true, fwPacket)
|
||||
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
|
||||
}
|
||||
|
||||
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
||||
dropReason := f.firewall.Drop(rxc.fwPacket.Packet, true, hostinfo, f.pki.GetCAPool(), rxc.ctCache.Get())
|
||||
if dropReason != nil {
|
||||
// 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. With UDP GRO this is a single segment of a shared
|
||||
// recvmmsg row whose capacity runs to the end of the whole row, so cap it to its own length (cap==len) to
|
||||
// keep the reject builder from writing past this segment into the next, not-yet-processed coalesced segment.
|
||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet[:len(packet):len(packet)], q)
|
||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, rxc.nb, rxc.scratch, rxc.q)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
||||
"fwPacket", fwPacket,
|
||||
"reason", dropReason,
|
||||
)
|
||||
hostinfo.logger(f.l).Debug("dropping inbound packet", "fwPacket", rxc.fwPacket, "reason", dropReason)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
err = f.batchers[q].Commit(out)
|
||||
err = f.batchers[rxc.q].Commit(out, batch.SortKey{Epoch: hostinfo.ConnectionState.epoch, Counter: messageCounter}, rxc.fwPacket)
|
||||
if err != nil {
|
||||
f.l.Error("Failed to write to tun", "error", err)
|
||||
}
|
||||
|
||||
+190
-24
@@ -9,15 +9,17 @@ import (
|
||||
|
||||
"github.com/google/gopacket"
|
||||
"github.com/google/gopacket/layers"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/net/ipv4"
|
||||
"golang.org/x/net/ipv6"
|
||||
)
|
||||
|
||||
func Test_newPacket(t *testing.T) {
|
||||
p := &firewall.Packet{}
|
||||
p := &firewall.ParsedPacket{}
|
||||
|
||||
// length fails
|
||||
err := newPacket([]byte{}, true, p)
|
||||
@@ -57,7 +59,7 @@ func Test_newPacket(t *testing.T) {
|
||||
Src: net.IPv4(10, 0, 0, 1),
|
||||
Dst: net.IPv4(10, 0, 0, 2),
|
||||
Options: []byte{0, 1, 0, 2},
|
||||
Protocol: firewall.ProtoTCP,
|
||||
Protocol: iputil.IPProtocolTCP,
|
||||
}
|
||||
|
||||
b, _ = h.Marshal()
|
||||
@@ -65,7 +67,7 @@ func Test_newPacket(t *testing.T) {
|
||||
err = newPacket(b, true, p)
|
||||
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
||||
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, uint16(3), p.RemotePort)
|
||||
@@ -96,7 +98,7 @@ func Test_newPacket(t *testing.T) {
|
||||
}
|
||||
|
||||
func Test_newPacket_v6(t *testing.T) {
|
||||
p := &firewall.Packet{}
|
||||
p := &firewall.ParsedPacket{}
|
||||
|
||||
// invalid ipv6
|
||||
ip := layers.IPv6{
|
||||
@@ -115,12 +117,12 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
err = newPacket(buffer.Bytes(), true, p)
|
||||
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||
|
||||
// A v6 packet with a hop-by-hop extension
|
||||
// ICMPv6 Payload (Echo Request)
|
||||
icmpLayer := layers.ICMPv6{
|
||||
TypeCode: layers.ICMPv6TypeEchoRequest,
|
||||
TypeCode: layers.CreateICMPv6TypeCode(layers.ICMPv6TypeEchoRequest, 0),
|
||||
}
|
||||
// Hop-by-Hop Extension Header
|
||||
hopOption := layers.IPv6HopByHopOption{}
|
||||
@@ -149,12 +151,12 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
// A full IPv6 header and 1 byte in the first extension, but missing
|
||||
// the length byte.
|
||||
err = newPacket(buffer.Bytes()[:41], true, p)
|
||||
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||
|
||||
// A full IPv6 header plus 1 full extension, but only 1 byte of the
|
||||
// next layer, missing length byte
|
||||
err = newPacket(buffer.Bytes()[:49], true, p)
|
||||
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||
err = nil
|
||||
|
||||
// A good ICMP packet
|
||||
@@ -167,7 +169,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
}
|
||||
|
||||
icmp := layers.ICMPv6{
|
||||
TypeCode: layers.ICMPv6TypeEchoRequest,
|
||||
TypeCode: layers.CreateICMPv6TypeCode(layers.ICMPv6TypeEchoRequest, 0),
|
||||
Checksum: 0x1234,
|
||||
}
|
||||
|
||||
@@ -189,6 +191,18 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
assert.Equal(t, uint16(0), p.LocalPort)
|
||||
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
|
||||
b := buffer.Bytes()
|
||||
b[6] = byte(layers.IPProtocolESP)
|
||||
@@ -213,16 +227,20 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
assert.Equal(t, uint16(0), p.LocalPort)
|
||||
assert.False(t, p.Fragment)
|
||||
|
||||
// An unknown protocol packet
|
||||
// An unknown protocol packet, we don't dissect it so we fail closed on its true protocol with no ports
|
||||
b = buffer.Bytes()
|
||||
b[6] = 255 // 255 is a reserved protocol number
|
||||
err = newPacket(b, true, p)
|
||||
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||
require.NoError(t, err)
|
||||
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
|
||||
ip = layers.IPv6{
|
||||
Version: 6,
|
||||
NextHeader: firewall.ProtoUDP,
|
||||
NextHeader: iputil.IPProtocolUDP,
|
||||
HopLimit: 128,
|
||||
SrcIP: net.IPv6linklocalallrouters,
|
||||
DstIP: net.IPv6linklocalallnodes,
|
||||
@@ -245,7 +263,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
// incoming
|
||||
err = newPacket(b, true, p)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
||||
assert.Equal(t, uint16(36123), p.RemotePort)
|
||||
@@ -255,7 +273,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
// outgoing
|
||||
err = newPacket(b, false, p)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
|
||||
assert.Equal(t, uint16(36123), p.LocalPort)
|
||||
@@ -272,7 +290,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
// incoming
|
||||
err = newPacket(b, true, p)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
||||
assert.Equal(t, uint16(36123), p.RemotePort)
|
||||
@@ -282,7 +300,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
// outgoing
|
||||
err = newPacket(b, false, p)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
|
||||
assert.Equal(t, uint16(36123), p.LocalPort)
|
||||
@@ -327,25 +345,25 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
|
||||
err = newPacket(b, true, p)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
||||
assert.Equal(t, uint16(36123), p.RemotePort)
|
||||
assert.Equal(t, uint16(22), p.LocalPort)
|
||||
assert.False(t, p.Fragment)
|
||||
|
||||
// Ensure buffer bounds checking during processing
|
||||
// Ensure buffer bounds checking during processing, a truncated AH header can't reach the payload
|
||||
err = newPacket(b[:41], true, p)
|
||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||
|
||||
// Invalid AH header
|
||||
b = buffer.Bytes()
|
||||
err = newPacket(b, true, p)
|
||||
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||
}
|
||||
|
||||
func Test_newPacket_ipv6Fragment(t *testing.T) {
|
||||
p := &firewall.Packet{}
|
||||
p := &firewall.ParsedPacket{}
|
||||
|
||||
ip := &layers.IPv6{
|
||||
Version: 6,
|
||||
@@ -525,7 +543,7 @@ func BenchmarkParseV6(b *testing.B) {
|
||||
secondFrag = append(secondFrag, fragHeader...)
|
||||
secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
|
||||
|
||||
fp := &firewall.Packet{}
|
||||
fp := &firewall.ParsedPacket{}
|
||||
|
||||
b.Run("Normal", func(b *testing.B) {
|
||||
for i := 0; i < b.N; i++ {
|
||||
@@ -649,7 +667,7 @@ func serializeAH(ah *layers.IPSecAH) []byte {
|
||||
// host OS parses the real header, a firewall port/proto bypass. The fix makes parseV6 land
|
||||
// on the same offset the host does.
|
||||
func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
||||
p := &firewall.Packet{}
|
||||
p := &firewall.ParsedPacket{}
|
||||
|
||||
const (
|
||||
hdrLen = 40 // IPv6 header
|
||||
@@ -661,7 +679,7 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
||||
pkt := make([]byte, realTCPAt+4)
|
||||
pkt[0] = 0x60 // version 6
|
||||
pkt[6] = byte(layers.IPProtocolIPv6Destination) // NextHeader -> Destination Options
|
||||
pkt[40] = byte(firewall.ProtoTCP) // Dest-Options NextHeader -> TCP
|
||||
pkt[40] = byte(iputil.IPProtocolTCP) // Dest-Options NextHeader -> TCP
|
||||
pkt[41] = 255 // HdrExtLen = 255
|
||||
|
||||
// Forged transport header at the pre-fix (wrong) offset: dst port 443.
|
||||
@@ -670,8 +688,156 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
||||
binary.BigEndian.PutUint16(pkt[realTCPAt+2:realTCPAt+4], 22)
|
||||
|
||||
require.NoError(t, newPacket(pkt, true, p))
|
||||
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
||||
// LocalPort is the destination port for incoming traffic. It must be the real port (22)
|
||||
// the host delivers to, not the forged 443 at the overflowed offset.
|
||||
assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset")
|
||||
}
|
||||
|
||||
// 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,28 +0,0 @@
|
||||
package batch
|
||||
|
||||
import "net/netip"
|
||||
|
||||
type RxBatcher interface {
|
||||
// Reserve creates a pkt to borrow
|
||||
Reserve(sz int) []byte
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush
|
||||
Commit(pkt []byte) error
|
||||
// Flush emits every queued packet in arrival order.
|
||||
// Returns the first error observed; keeps draining so one bad packet doesn't hold up the rest.
|
||||
// After Flush returns, borrowed payload slices may be recycled.
|
||||
Flush() error
|
||||
}
|
||||
|
||||
type TxBatcher interface {
|
||||
// Reserve creates a pkt to borrow
|
||||
Reserve(sz int) []byte
|
||||
// Commit borrows pkt and records its destination plus the 2-bit
|
||||
// IP-level ECN codepoint to set on the outer (carrier) header. The
|
||||
// caller must keep pkt valid until the next Flush. Pass 0 (Not-ECT)
|
||||
// to leave the outer ECN field unset.
|
||||
Commit(pkt []byte, dst netip.AddrPort, outerECN byte)
|
||||
// Flush emits every queued packet via the underlying batch writer in arrival order.
|
||||
// Returns an errors.Join of one or more errors.
|
||||
// After Flush returns, borrowed payload slices may be recycled.
|
||||
Flush() error
|
||||
}
|
||||
@@ -0,0 +1,187 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
+114
-126
@@ -5,138 +5,136 @@ import (
|
||||
"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.
|
||||
// openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto never alias.
|
||||
type flowKey struct {
|
||||
src, dst [16]byte
|
||||
sport, dport uint16
|
||||
isV6 bool
|
||||
}
|
||||
|
||||
// initialSlots is the starting capacity of the slot pool. One flow per
|
||||
// packet is the worst case so this matches a typical carrier-side
|
||||
// recvmmsg batch on the encrypted UDP socket.
|
||||
// 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
|
||||
|
||||
// parsedIP is the IP-level result of parseIPPrologue. The caller layers
|
||||
// L4-specific parsing (TCP / UDP) on top.
|
||||
type parsedIP struct {
|
||||
fk flowKey
|
||||
ipHdrLen int
|
||||
// pkt is the original buffer trimmed to the IP-declared total length.
|
||||
// Anything below the IP layer (transport parsers) should slice into
|
||||
// pkt rather than the unbounded original.
|
||||
pkt []byte
|
||||
// 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
|
||||
}
|
||||
|
||||
// parseIPPrologue extracts the IP-level fields the coalescers care about:
|
||||
// IHL/payload length, version, src/dst addresses, and the L4 protocol byte.
|
||||
// Returns ok=false for malformed input, IPv4 with options or fragmentation,
|
||||
// or IPv6 with extension headers (all rejected by both coalescers in
|
||||
// identical ways before this refactor).
|
||||
//
|
||||
// On success, p.pkt is len-trimmed to the IP-declared length so callers
|
||||
// don't have to repeat the trim. wantProto is the IANA protocol number to
|
||||
// require (6 for TCP, 17 for UDP); ok=false for any other value.
|
||||
func parseIPPrologue(pkt []byte, wantProto byte) (parsedIP, bool) {
|
||||
var p parsedIP
|
||||
if len(pkt) < 20 {
|
||||
return p, false
|
||||
// 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
|
||||
}
|
||||
v := pkt[0] >> 4
|
||||
switch v {
|
||||
case 4:
|
||||
ihl := int(pkt[0]&0x0f) * 4
|
||||
if ihl != 20 {
|
||||
return p, false
|
||||
}
|
||||
if pkt[9] != wantProto {
|
||||
return p, false
|
||||
}
|
||||
// Reject actual fragmentation (MF or non-zero frag offset).
|
||||
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
||||
return p, false
|
||||
}
|
||||
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
|
||||
if totalLen > len(pkt) || totalLen < ihl {
|
||||
return p, false
|
||||
}
|
||||
p.ipHdrLen = 20
|
||||
p.fk.isV6 = false
|
||||
copy(p.fk.src[:4], pkt[12:16])
|
||||
copy(p.fk.dst[:4], pkt[16:20])
|
||||
p.pkt = pkt[:totalLen]
|
||||
case 6:
|
||||
if len(pkt) < 40 {
|
||||
return p, false
|
||||
}
|
||||
if pkt[6] != wantProto {
|
||||
return p, false
|
||||
}
|
||||
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
|
||||
if 40+payloadLen > len(pkt) {
|
||||
return p, false
|
||||
}
|
||||
p.ipHdrLen = 40
|
||||
p.fk.isV6 = true
|
||||
copy(p.fk.src[:], pkt[8:24])
|
||||
copy(p.fk.dst[:], pkt[24:40])
|
||||
p.pkt = pkt[:40+payloadLen]
|
||||
default:
|
||||
return p, false
|
||||
// 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
|
||||
}
|
||||
return p, true
|
||||
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.
|
||||
// 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.
|
||||
// The transport (L4) portion of the header is checked separately by the per-protocol matcher.
|
||||
func ipHeadersMatch(a, b []byte, isV6 bool) bool {
|
||||
if isV6 {
|
||||
// IPv6: byte 0 = version/TC[7:4], byte 1 = TC[3:0]/flow[19:16],
|
||||
// bytes [2:4] = flow[15:0], [6:8] = next_hdr/hop, [8:40] = src+dst.
|
||||
// Compare byte 1 fully so ECN (TC[1:0]) must match. Skip [4:6] payload_len.
|
||||
if a[0] != b[0] {
|
||||
return false
|
||||
}
|
||||
if a[1] != b[1] {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[2:4], b[2:4]) {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[6:40], b[6:40]) {
|
||||
return false
|
||||
}
|
||||
// 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
|
||||
}
|
||||
// IPv4: byte 0 = version/IHL, byte 1 = DSCP(6)|ECN(2),
|
||||
// [6:10] flags/fragoff/TTL/proto, [12:20] src+dst.
|
||||
// Compare byte 1 fully so ECN must match.
|
||||
// Skip [2:4] total len, [4:6] id, [10:12] csum.
|
||||
if a[0] != b[0] {
|
||||
return false
|
||||
}
|
||||
if a[1] != b[1] {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[6:10], b[6:10]) {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[12:20], b[12:20]) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
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
|
||||
@@ -145,17 +143,14 @@ type Arena struct {
|
||||
buf []byte
|
||||
}
|
||||
|
||||
// NewArena returns an Arena with a pre-allocated backing of the given
|
||||
// capacity. Pass 0 if you don't intend to call Reserve (e.g. a test that
|
||||
// only feeds the coalescer pre-made []byte packets via Commit).
|
||||
// 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.
|
||||
// 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)
|
||||
@@ -166,16 +161,9 @@ func (a *Arena) Reserve(sz int) []byte {
|
||||
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.
|
||||
// Reset releases every slice handed out since the last Reset.
|
||||
// Callers must not use any previously-borrowed slice after this returns.
|
||||
// The underlying backing array is retained so subsequent Reserves don't re-allocate.
|
||||
func (a *Arena) Reset() {
|
||||
a.buf = a.buf[:0]
|
||||
}
|
||||
|
||||
// Reserver hands out an sz-byte slice valid until its Resetter runs.
|
||||
type Reserver func(sz int) []byte
|
||||
|
||||
// Resetter clears all reservations held by a Reserver. Only the arena's
|
||||
// owner holds one; lanes inside a MultiCoalescer get nil.
|
||||
type Resetter func()
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
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))
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
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,119 +1,121 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"slices"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
)
|
||||
|
||||
// MultiCoalescer fans plaintext packets out to lane-specific batchers based
|
||||
// on the IP/L4 protocol of the packet, sharing a single Reserve arena
|
||||
// across lanes so the caller's allocation pattern is unchanged.
|
||||
// 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.
|
||||
//
|
||||
// Lanes are processed independently: the TCP coalescer only sees TCP, the
|
||||
// UDP coalescer only sees UDP, and the passthrough lane handles everything
|
||||
// else. Per-flow arrival order is preserved because a single 5-tuple only
|
||||
// ever lands in one lane and each lane preserves its own slot order.
|
||||
// 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.
|
||||
//
|
||||
// Cross-lane order is NOT preserved across the TCP/UDP/passthrough split.
|
||||
// This is acceptable because the carrier-side recvmmsg path already
|
||||
// stable-sorts by (peer, message counter) before delivering plaintext
|
||||
// here, so replay-window invariants are unaffected, and apps observe
|
||||
// correct per-flow ordering — which is all the IP layer guarantees anyway.
|
||||
// Do not "fix" this by interleaving lane outputs at flush time; that
|
||||
// negates the entire point of coalescing (each lane needs to see runs of
|
||||
// adjacent same-flow packets to coalesce them).
|
||||
// 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
|
||||
// arena is owned by the Multi: lanes get only its Reserve (nil Resetter)
|
||||
// and Flush resets it exactly once after every lane has drained.
|
||||
arena *Arena
|
||||
|
||||
// staged holds this batch's packets and sort keys until Flush. Borrowed: the caller keeps
|
||||
// each pkt alive until Flush returns.
|
||||
staged []stagedPacket
|
||||
}
|
||||
|
||||
// DefaultMultiArenaCap is the recommended arena capacity for a Multi-lane
|
||||
// batcher: 64 slots × 65535 bytes ≈ 4 MiB, enough to hold one recvmmsg
|
||||
// burst worth of MTU-sized packets without the arena growing.
|
||||
const DefaultMultiArenaCap = initialSlots * 65535
|
||||
// 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. tcpEnabled lets the caller
|
||||
// opt out of TCP coalescing (e.g. when the queue can't do TSO); udpEnabled
|
||||
// likewise gates UDP coalescing (only enable when USO was negotiated).
|
||||
// Either lane disabled redirects its traffic into the passthrough lane.
|
||||
// arena is the single backing slab shared across every lane; the caller
|
||||
// pre-sizes it via NewArena so the hot path never allocates.
|
||||
func NewMultiCoalescer(w io.Writer, l *slog.Logger, arena *Arena, tcpEnabled, udpEnabled bool) *MultiCoalescer {
|
||||
// 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, arena.Reserve, nil),
|
||||
arena: arena,
|
||||
}
|
||||
if tcpEnabled {
|
||||
m.tcp = NewTCPCoalescer(w, l, arena.Reserve, nil)
|
||||
}
|
||||
if udpEnabled {
|
||||
m.udp = NewUDPCoalescer(w, arena.Reserve, nil)
|
||||
pt: NewPassthrough(w),
|
||||
staged: make([]stagedPacket, 0, initialSlots),
|
||||
}
|
||||
m.tcp = NewTCPCoalescer(w, l)
|
||||
m.udp = NewUDPCoalescer(w)
|
||||
return m
|
||||
}
|
||||
|
||||
func (m *MultiCoalescer) Reserve(sz int) []byte {
|
||||
return m.arena.Reserve(sz)
|
||||
// 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
|
||||
}
|
||||
|
||||
// Commit dispatches pkt to the appropriate lane based on IP version + L4
|
||||
// proto. Borrowed slice contract is identical to the single-lane batchers,
|
||||
// pkt must remain valid until the next Flush.
|
||||
//
|
||||
// On the success path the IP/TCP-or-UDP parse happens here once and the
|
||||
// parsed struct is handed to the lane via commitParsed so the lane doesn't
|
||||
// re-walk the header.
|
||||
func (m *MultiCoalescer) Commit(pkt []byte) error {
|
||||
if len(pkt) < 20 {
|
||||
return m.pt.Commit(pkt)
|
||||
// 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
|
||||
}
|
||||
v := pkt[0] >> 4
|
||||
var proto byte
|
||||
switch v {
|
||||
case 4:
|
||||
proto = pkt[9]
|
||||
case 6:
|
||||
if len(pkt) < 40 {
|
||||
return m.pt.Commit(pkt)
|
||||
}
|
||||
proto = pkt[6]
|
||||
default:
|
||||
return m.pt.Commit(pkt)
|
||||
}
|
||||
switch proto {
|
||||
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 {
|
||||
info, ok := parseTCPBase(pkt)
|
||||
if !ok {
|
||||
// Malformed/unsupported TCP shape (IP options, fragments, ...).
|
||||
// Handle this via passthrough support in the TCP coalescer, to attempt to preserve flow order.
|
||||
m.tcp.addPassthrough(pkt)
|
||||
return nil
|
||||
}
|
||||
return m.tcp.commitParsed(pkt, info)
|
||||
return m.tcp.commitStaged(sp)
|
||||
}
|
||||
case ipProtoUDP:
|
||||
if m.udp != nil {
|
||||
info, ok := parseUDP(pkt)
|
||||
if !ok {
|
||||
m.udp.addPassthrough(pkt) //we could also m.pt.Commit() here I guess?
|
||||
return nil
|
||||
}
|
||||
return m.udp.commitParsed(pkt, info)
|
||||
return m.udp.commitStaged(sp)
|
||||
}
|
||||
}
|
||||
return m.pt.Commit(pkt)
|
||||
return m.pt.enqueue(sp.pkt)
|
||||
}
|
||||
|
||||
// Flush drains every lane in a fixed order, then resets the shared arena once.
|
||||
// A lane error doesn't stop the remaining lanes; the joined errors are returned.
|
||||
// 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)
|
||||
@@ -127,6 +129,5 @@ func (m *MultiCoalescer) Flush() error {
|
||||
if err := m.pt.Flush(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
m.arena.Reset()
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
@@ -1,17 +1,39 @@
|
||||
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 := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, true)
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
|
||||
tcpPay := make([]byte, 1200)
|
||||
udpPay := make([]byte, 1200)
|
||||
@@ -21,19 +43,19 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
||||
icmp[3] = 28
|
||||
icmp[9] = 1
|
||||
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay)); err != nil {
|
||||
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)); err != nil {
|
||||
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)); err != nil {
|
||||
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)); err != nil {
|
||||
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); err != nil {
|
||||
if err := m.Commit(icmp, k.next(), testPP(icmp)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
@@ -48,17 +70,162 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerDisabledUDPFallsThrough verifies that when the UDP lane
|
||||
// is disabled (e.g. kernel doesn't support USO), UDP packets still reach
|
||||
// the kernel via the passthrough lane rather than being lost.
|
||||
func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
|
||||
// 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 := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, false) // TSO on, USO off
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||
// 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(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||
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 {
|
||||
@@ -72,16 +239,164 @@ func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerDisabledTCPFallsThrough mirrors the TSO=off case.
|
||||
func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), false, true) // TSO off, USO on
|
||||
|
||||
pay := make([]byte, 1200)
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
||||
// 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 err := m.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil {
|
||||
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 {
|
||||
@@ -94,3 +409,29 @@ func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -2,54 +2,29 @@ package batch
|
||||
|
||||
import (
|
||||
"io"
|
||||
|
||||
"github.com/slackhq/nebula/udp"
|
||||
)
|
||||
|
||||
// Passthrough is a RxBatcher that doesn't batch anything, it just accumulates and then sends packets.
|
||||
// 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
|
||||
reserver Reserver
|
||||
resetter Resetter
|
||||
cursor int
|
||||
out io.Writer
|
||||
slots [][]byte
|
||||
}
|
||||
|
||||
const passthroughBaseNumSlots = 128
|
||||
|
||||
// DefaultPassthroughArenaCap is the recommended arena capacity for a
|
||||
// standalone Passthrough batcher: 128 slots × udp.MTU ≈ 1.1 MiB.
|
||||
const DefaultPassthroughArenaCap = passthroughBaseNumSlots * udp.MTU
|
||||
|
||||
func NewPassthrough(w io.Writer, reserver Reserver, resetter Resetter) *Passthrough {
|
||||
func NewPassthrough(w io.Writer) *Passthrough {
|
||||
return &Passthrough{
|
||||
out: w,
|
||||
slots: make([][]byte, 0, passthroughBaseNumSlots),
|
||||
reserver: reserver,
|
||||
resetter: resetter,
|
||||
out: w,
|
||||
slots: make([][]byte, 0, 128),
|
||||
}
|
||||
}
|
||||
|
||||
func (p *Passthrough) Reserve(sz int) []byte {
|
||||
return p.reserver(sz)
|
||||
}
|
||||
|
||||
func (p *Passthrough) Commit(pkt []byte) error {
|
||||
// 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
|
||||
}
|
||||
|
||||
// Flush drains every queued packet and calls the configured Resetter
|
||||
func (p *Passthrough) Flush() error {
|
||||
firstErr := p.drain()
|
||||
if p.resetter != nil {
|
||||
p.resetter()
|
||||
}
|
||||
return firstErr
|
||||
}
|
||||
|
||||
// drain writes out every queued packet and clears the slot list.
|
||||
func (p *Passthrough) drain() error {
|
||||
var firstErr error
|
||||
for _, s := range p.slots {
|
||||
_, err := p.out.Write(s)
|
||||
|
||||
+182
-438
@@ -2,12 +2,9 @@ package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"slices"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
)
|
||||
@@ -23,24 +20,19 @@ const tcpCoalesceBufSize = 65535
|
||||
// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency.
|
||||
const tcpCoalesceMaxSegs = 64
|
||||
|
||||
// tcpCoalesceHdrCap is the scratch space we copy a seed's IP+TCP header
|
||||
// into. IPv6 (40) + TCP with full options (60) = 100 bytes.
|
||||
const tcpCoalesceHdrCap = 100
|
||||
|
||||
// coalesceSlot is one entry in the coalescer's ordered event queue. When
|
||||
// passthrough is true the slot holds a single borrowed packet that must be
|
||||
// emitted verbatim (non-TCP, non-admissible TCP, or oversize seed). When
|
||||
// passthrough is false the slot is an in-progress coalesced superpacket:
|
||||
// hdrBuf is a mutable copy of the seed's IP+TCP header (we patch total
|
||||
// length and pseudo-header partial at flush), and payIovs are *borrowed*
|
||||
// slices from the caller's plaintext buffers — no payload is ever copied.
|
||||
// The caller (listenOut) must keep those buffers alive until Flush.
|
||||
// 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 {
|
||||
passthrough bool
|
||||
rawPkt []byte // borrowed when passthrough
|
||||
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
|
||||
hdrBuf [tcpCoalesceHdrCap]byte
|
||||
hdrLen int
|
||||
ipHdrLen int
|
||||
isV6 bool
|
||||
@@ -48,216 +40,203 @@ type coalesceSlot struct {
|
||||
numSeg int
|
||||
totalPay int
|
||||
nextSeq uint32
|
||||
// psh closes the chain: set when the last-accepted segment had PSH or
|
||||
// was sub-gsoSize. No further appends after that.
|
||||
psh bool
|
||||
payIovs [][]byte
|
||||
payIovs [][]byte
|
||||
}
|
||||
|
||||
// TCPCoalescer accumulates adjacent in-flow TCP data segments across
|
||||
// multiple concurrent flows and emits each flow's run as a single TSO
|
||||
// superpacket via tio.GSOWriter. All output — coalesced or not — is
|
||||
// deferred until Flush so arrival order is preserved on the wire. Owns
|
||||
// no locks; one coalescer per TUN write queue.
|
||||
// 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 {
|
||||
plainW io.Writer
|
||||
gsoW tio.GSOWriter // nil when the queue doesn't support TSO
|
||||
w tio.GSOWriter
|
||||
|
||||
// slots is the ordered event queue. Flush walks it once and emits each
|
||||
// entry as either a WriteGSO (coalesced) or a plainW.Write (passthrough).
|
||||
// entry as either a WriteGSO (coalesced) or a w.Write (verbatim).
|
||||
slots []*coalesceSlot
|
||||
// openSlots maps a flow key to its most recent non-sealed slot, so new
|
||||
// segments can extend an in-progress superpacket in O(1). Slots are
|
||||
// removed from this map when they close (PSH or short-last-segment),
|
||||
// when a non-admissible packet for that flow arrives, or in Flush.
|
||||
// openSlots 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. Steady-state
|
||||
// bulk traffic is dominated by a single flow, so comparing the
|
||||
// incoming key against the cached slot's own fk lets the hot path
|
||||
// skip the map lookup (and the aeshash of a 38-byte key) entirely.
|
||||
// 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/sealed.
|
||||
// at is removed.
|
||||
lastSlot *coalesceSlot
|
||||
pool []*coalesceSlot // free list for reuse
|
||||
reserver Reserver
|
||||
resetter Resetter
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
func NewTCPCoalescer(w io.Writer, l *slog.Logger, reserver Reserver, resetter Resetter) *TCPCoalescer {
|
||||
c := &TCPCoalescer{
|
||||
plainW: w,
|
||||
// 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),
|
||||
reserver: reserver,
|
||||
resetter: resetter,
|
||||
l: l,
|
||||
}
|
||||
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP); ok {
|
||||
c.gsoW = gw
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// parsedTCP holds the fields extracted from a single parse so later steps
|
||||
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||
type parsedTCP struct {
|
||||
fk flowKey
|
||||
ipHdrLen int
|
||||
tcpHdrLen int
|
||||
hdrLen int
|
||||
payLen int
|
||||
seq uint32
|
||||
flags byte
|
||||
fk flowKey
|
||||
ipHdrLen int
|
||||
hdrLen int
|
||||
payLen int
|
||||
seq uint32
|
||||
flags byte
|
||||
}
|
||||
|
||||
// parseTCPBase extracts the flow key and IP/TCP offsets for any TCP packet,
|
||||
// regardless of whether it's admissible for coalescing. Returns ok=false
|
||||
// for non-TCP or malformed input.
|
||||
// Accepts IPv4 (no options or fragmentation) and IPv6 (no extension headers).
|
||||
func parseTCPBase(pkt []byte) (parsedTCP, bool) {
|
||||
var p parsedTCP
|
||||
ip, ok := parseIPPrologue(pkt, ipProtoTCP)
|
||||
// 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 p, false
|
||||
return false
|
||||
}
|
||||
pkt = ip.pkt
|
||||
p.fk = ip.fk
|
||||
p.ipHdrLen = ip.ipHdrLen
|
||||
|
||||
if len(pkt) < p.ipHdrLen+20 {
|
||||
return p, false
|
||||
}
|
||||
tcpOff := int(pkt[p.ipHdrLen+12]>>4) * 4
|
||||
if tcpOff < 20 || tcpOff > 60 {
|
||||
return p, false
|
||||
}
|
||||
if len(pkt) < p.ipHdrLen+tcpOff {
|
||||
return p, false
|
||||
}
|
||||
p.tcpHdrLen = tcpOff
|
||||
p.hdrLen = p.ipHdrLen + tcpOff
|
||||
p.payLen = len(pkt) - p.hdrLen
|
||||
p.seq = binary.BigEndian.Uint32(pkt[p.ipHdrLen+4 : p.ipHdrLen+8])
|
||||
p.flags = pkt[p.ipHdrLen+13]
|
||||
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
||||
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
||||
return p, true
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
// TCP flag bits (byte 13 of the TCP header). Only the bits actually consulted
|
||||
// by the coalescer are named; FIN/SYN/RST/URG/CWR are rejected via the
|
||||
// negative mask in coalesceable, not by name.
|
||||
// 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
|
||||
)
|
||||
|
||||
// coalesceable reports whether a parsed TCP segment is eligible for
|
||||
// coalescing. Accepts ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE with a
|
||||
// non-empty payload. CWR is excluded because it marks a one-shot
|
||||
// congestion-window-reduced transition the receiver must observe at a
|
||||
// segment boundary.
|
||||
func (p parsedTCP) coalesceable() bool {
|
||||
if p.flags&tcpFlagAck == 0 {
|
||||
return false
|
||||
}
|
||||
if p.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
|
||||
return false
|
||||
}
|
||||
return p.payLen > 0
|
||||
// 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
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) Reserve(sz int) []byte {
|
||||
return c.reserver(sz)
|
||||
// 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)
|
||||
}
|
||||
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||
func (c *TCPCoalescer) Commit(pkt []byte) error {
|
||||
if c.gsoW == nil {
|
||||
c.addPassthrough(pkt)
|
||||
// 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
|
||||
}
|
||||
info, ok := parseTCPBase(pkt)
|
||||
if !ok {
|
||||
c.addPassthrough(pkt)
|
||||
var info parsedTCP
|
||||
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(pkt, info)
|
||||
return c.commitParsed(sp.pkt, &info)
|
||||
}
|
||||
|
||||
// commitParsed is the post-parse half of Commit. The caller must have
|
||||
// already verified parseTCPBase succeeded (info is a valid TCP parse).
|
||||
// Used by MultiCoalescer.Commit to avoid re-walking the IP/TCP header
|
||||
// after the dispatcher has already done so.
|
||||
func (c *TCPCoalescer) commitParsed(pkt []byte, info parsedTCP) error {
|
||||
if c.gsoW == nil {
|
||||
c.addPassthrough(pkt)
|
||||
// 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.coalesceable() {
|
||||
// TCP but not admissible (SYN/FIN/RST/URG/CWR or zero-payload).
|
||||
// Seal this flow's open slot so later in-flow packets don't extend
|
||||
// it and accidentally reorder past this passthrough.
|
||||
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||
c.lastSlot = nil
|
||||
}
|
||||
delete(c.openSlots, info.fk)
|
||||
c.addPassthrough(pkt)
|
||||
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
|
||||
}
|
||||
|
||||
// Single-flow fast path: with only one open flow the cache hits every
|
||||
// packet, and len(openSlots)==1 lets us skip the 38-byte fk compare
|
||||
// when there are multiple flows in flight (where the hit rate would
|
||||
// be ~0 and the compare is pure overhead).
|
||||
// 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 && len(c.openSlots) == 1 && last.fk == info.fk {
|
||||
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) {
|
||||
c.appendPayload(open, pkt, info)
|
||||
if open.psh {
|
||||
delete(c.openSlots, info.fk)
|
||||
c.lastSlot = nil
|
||||
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 — seal it and fall through to seed a fresh slot.
|
||||
delete(c.openSlots, info.fk)
|
||||
if c.lastSlot == open {
|
||||
c.lastSlot = 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
|
||||
}
|
||||
|
||||
// Flush emits every queued event in (per-flow) seq order.
|
||||
func (c *TCPCoalescer) Flush() error {
|
||||
first := c.drain()
|
||||
if c.resetter != nil {
|
||||
c.resetter()
|
||||
}
|
||||
return first
|
||||
}
|
||||
|
||||
// drain emits every queued slot (reordering/merging coalesced runs first)
|
||||
// and clears the slot state.
|
||||
func (c *TCPCoalescer) drain() error {
|
||||
c.reorderForFlush()
|
||||
var first error
|
||||
for _, s := range c.slots {
|
||||
var err error
|
||||
if s.passthrough {
|
||||
_, err = c.plainW.Write(s.rawPkt)
|
||||
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)
|
||||
}
|
||||
@@ -274,23 +253,27 @@ func (c *TCPCoalescer) drain() error {
|
||||
return first
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) addPassthrough(pkt []byte) {
|
||||
func (c *TCPCoalescer) addVerbatim(pkt []byte) {
|
||||
s := c.take()
|
||||
s.passthrough = true
|
||||
s.verbatim = true
|
||||
s.rawPkt = pkt
|
||||
c.slots = append(c.slots, s)
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
|
||||
if info.hdrLen > tcpCoalesceHdrCap || info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
||||
// Pathological shape — can't fit our scratch, emit as-is.
|
||||
c.addPassthrough(pkt)
|
||||
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.passthrough = false
|
||||
s.rawPkt = nil
|
||||
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
||||
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
|
||||
@@ -299,26 +282,23 @@ func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
|
||||
s.numSeg = 1
|
||||
s.totalPay = info.payLen
|
||||
s.nextSeq = info.seq + uint32(info.payLen)
|
||||
s.psh = info.flags&tcpFlagPsh != 0
|
||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
c.slots = append(c.slots, s)
|
||||
if !s.psh {
|
||||
if info.flags&tcpFlagPsh == 0 {
|
||||
c.openSlots[info.fk] = s
|
||||
c.lastSlot = s
|
||||
} else if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||
// PSH-on-seed seals the slot immediately. Any prior cached open
|
||||
// slot for this flow has just been sealed-and-replaced by this
|
||||
// passthrough-shaped seed, so drop the cache too.
|
||||
c.lastSlot = nil
|
||||
} 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, chain not closed.
|
||||
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bool {
|
||||
if s.psh {
|
||||
return false
|
||||
}
|
||||
// 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
|
||||
}
|
||||
@@ -334,31 +314,35 @@ func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bo
|
||||
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
||||
return false
|
||||
}
|
||||
// ECE state must be stable across a burst — receivers expect the
|
||||
// flag set on every segment of a CE-echoing window or none.
|
||||
seedFlags := s.hdrBuf[s.ipHdrLen+13]
|
||||
// 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 !headersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||
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
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info parsedTCP) {
|
||||
// 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. Without this the sender's push signal is dropped.
|
||||
s.hdrBuf[s.ipHdrLen+13] |= tcpFlagPsh
|
||||
}
|
||||
if info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0 {
|
||||
s.psh = true
|
||||
// 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 {
|
||||
@@ -372,21 +356,18 @@ func (c *TCPCoalescer) take() *coalesceSlot {
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
||||
s.passthrough = false
|
||||
s.rawPkt = nil
|
||||
clear(s.payIovs)
|
||||
s.payIovs = s.payIovs[:0]
|
||||
s.numSeg = 0
|
||||
s.totalPay = 0
|
||||
s.psh = false
|
||||
*s = coalesceSlot{payIovs: s.payIovs[:0]}
|
||||
c.pool = append(c.pool, s)
|
||||
}
|
||||
|
||||
// flushSlot patches the header and calls WriteGSO. Does not remove the slot from c.slots.
|
||||
// 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.hdrBuf[:s.hdrLen]
|
||||
hdr := s.rawPkt[:s.hdrLen]
|
||||
|
||||
if s.isV6 {
|
||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
||||
@@ -406,7 +387,7 @@ func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
||||
tcsum := s.ipHdrLen + 16
|
||||
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
||||
|
||||
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
|
||||
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
|
||||
@@ -437,242 +418,6 @@ func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// reorderForFlush neutralizes wire-side reorder that the rxOrder buffer
|
||||
// couldn't catch (anything crossing a recvmmsg batch boundary). Without
|
||||
// this pass a small wire reorder — counter 250 arriving in batch K when
|
||||
// 200..249 are coming in batch K+1 — would seed an out-of-seq slot first
|
||||
// and emit it ahead of the lower-seq slot, manifesting at the inner TCP
|
||||
// receiver as a much larger reorder than the wire actually had.
|
||||
//
|
||||
// Two phases:
|
||||
// 1. Sort each passthrough-bounded segment of c.slots by (flow, seq).
|
||||
// Cross-flow ordering inside a segment isn't preserved (it never was
|
||||
// and doesn't matter for any single flow's TCP correctness).
|
||||
// 2. Sweep once and merge adjacent same-flow slots whose ranges are now
|
||||
// contiguous AND whose tail is gsoSize-aligned. The tail constraint
|
||||
// matters because the kernel TSO splitter chops at gsoSize from the
|
||||
// start of the merged payload — a short segment in the middle would
|
||||
// desynchronize every later segment.
|
||||
//
|
||||
// Passthrough slots act as barriers: the merge check skips them on either
|
||||
// side, so a SYN/FIN/RST/CWR is never reordered relative to its flow's
|
||||
// data.
|
||||
func (c *TCPCoalescer) reorderForFlush() {
|
||||
if len(c.slots) <= 1 {
|
||||
return
|
||||
}
|
||||
runStart := 0
|
||||
for i := 0; i <= len(c.slots); i++ {
|
||||
if i < len(c.slots) && !c.slots[i].passthrough {
|
||||
continue
|
||||
}
|
||||
c.sortRun(c.slots[runStart:i])
|
||||
runStart = i + 1
|
||||
}
|
||||
out := c.slots[:0]
|
||||
logged := false
|
||||
for _, s := range c.slots {
|
||||
if n := len(out); n > 0 {
|
||||
prev := out[n-1]
|
||||
if !prev.passthrough && !s.passthrough && prev.fk == s.fk {
|
||||
// Same-flow neighbors after sort. If they aren't seq-
|
||||
// contiguous it's a real gap — packets the wire reordered
|
||||
// across batches, or actual loss before nebula. Log it so
|
||||
// the operator can quantify how often it happens; the data
|
||||
// itself still emits in seq order, kernel TCP handles the
|
||||
// gap via its OOO queue.
|
||||
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
if prev.nextSeq != slotSeedSeq(s) {
|
||||
logged = true
|
||||
gap := int64(slotSeedSeq(s)) - int64(prev.nextSeq)
|
||||
c.l.Debug("tcp coalesce: cross-slot seq gap",
|
||||
"src", flowKeyAddr(s.fk, false),
|
||||
"dst", flowKeyAddr(s.fk, true),
|
||||
"sport", s.fk.sport,
|
||||
"dport", s.fk.dport,
|
||||
"prev_seed_seq", slotSeedSeq(prev),
|
||||
"prev_next_seq", prev.nextSeq,
|
||||
"this_seed_seq", slotSeedSeq(s),
|
||||
"gap_bytes", gap,
|
||||
"prev_seg_count", prev.numSeg,
|
||||
"prev_total_pay", prev.totalPay,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
if canMergeSlots(prev, s) {
|
||||
mergeSlots(prev, s)
|
||||
c.release(s)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
out = append(out, s)
|
||||
}
|
||||
if logged {
|
||||
c.l.Warn("==== end of batch ====")
|
||||
}
|
||||
c.slots = out
|
||||
}
|
||||
|
||||
// flowKeyAddr returns the src or dst address from fk as a netip.Addr for
|
||||
// logging. Only used on the cold gap-log path so the netip allocation
|
||||
// doesn't matter.
|
||||
func flowKeyAddr(fk flowKey, dst bool) netip.Addr {
|
||||
src := fk.src
|
||||
if dst {
|
||||
src = fk.dst
|
||||
}
|
||||
if fk.isV6 {
|
||||
return netip.AddrFrom16(src)
|
||||
}
|
||||
var v4 [4]byte
|
||||
copy(v4[:], src[:4])
|
||||
return netip.AddrFrom4(v4)
|
||||
}
|
||||
|
||||
// sortRun stable-sorts run by (flowKey, seedSeq) so each flow's slots
|
||||
// cluster together in seq order, ready for the merge sweep. Stable so
|
||||
// equal-key slots keep their original relative position (defensive — a
|
||||
// duplicate seedSeq would already mean something's wrong upstream).
|
||||
func (c *TCPCoalescer) sortRun(run []*coalesceSlot) {
|
||||
if len(run) <= 1 {
|
||||
return
|
||||
}
|
||||
// slices.SortStableFunc with a free, non-capturing comparator avoids the
|
||||
// reflection + closure-escape allocations that sort.SliceStable forces.
|
||||
slices.SortStableFunc(run, compareCoalesceSlots)
|
||||
}
|
||||
|
||||
func compareCoalesceSlots(a, b *coalesceSlot) int {
|
||||
if cmp := flowKeyCompare(a.fk, b.fk); cmp != 0 {
|
||||
return cmp
|
||||
}
|
||||
aSeq, bSeq := slotSeedSeq(a), slotSeedSeq(b)
|
||||
if aSeq == bSeq {
|
||||
return 0
|
||||
}
|
||||
if tcpSeqLess(aSeq, bSeq) {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
|
||||
// slotSeedSeq returns the TCP seq of the slot's seed (first segment).
|
||||
// nextSeq tracks the seq just past the last appended byte; subtracting
|
||||
// totalPay walks back to the seed. uint32 wraparound is the right TCP
|
||||
// arithmetic so no special-casing is needed.
|
||||
func slotSeedSeq(s *coalesceSlot) uint32 {
|
||||
return s.nextSeq - uint32(s.totalPay)
|
||||
}
|
||||
|
||||
// tcpSeqLess reports whether a precedes b in TCP serial-number arithmetic
|
||||
// (RFC 1323 §2.3). The signed int32 cast turns the modular subtraction
|
||||
// into the right comparison even across the 2^32 wrap.
|
||||
func tcpSeqLess(a, b uint32) bool {
|
||||
return int32(a-b) < 0
|
||||
}
|
||||
|
||||
// flowKeyCompare orders flowKeys deterministically. The exact ordering
|
||||
// is irrelevant — only that same-flow slots cluster together so the
|
||||
// post-sort sweep can merge contiguous pairs.
|
||||
func flowKeyCompare(a, b flowKey) int {
|
||||
// Cheap scalar fields first so most non-matching keys short-circuit
|
||||
// without ever calling bytes.Compare. sport is the ephemeral port on
|
||||
// egress flows and discriminates fastest. For matching keys (same
|
||||
// flow), array equality on src/dst inlines to word-sized compares,
|
||||
// so we only pay bytes.Compare when the arrays actually differ.
|
||||
if a.sport != b.sport {
|
||||
if a.sport < b.sport {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
if a.dport != b.dport {
|
||||
if a.dport < b.dport {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
if a.dst != b.dst {
|
||||
return bytes.Compare(a.dst[:], b.dst[:])
|
||||
}
|
||||
if a.src != b.src {
|
||||
return bytes.Compare(a.src[:], b.src[:])
|
||||
}
|
||||
if a.isV6 != b.isV6 {
|
||||
if !a.isV6 {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// canMergeSlots reports whether s can fold into prev as one merged TSO
|
||||
// superpacket. Same flow, contiguous TCP byte range, equal gsoSize, and
|
||||
// fits within the kernel TSO limits. The tail-of-prev check rejects any
|
||||
// merge whose first slot ended on a sub-gsoSize segment — kernel TSO
|
||||
// would split the merged skb at gsoSize boundaries from the start, so a
|
||||
// short segment in the middle would corrupt every later segment. PSH and
|
||||
// ECE state must agree across both slots: PSH is a semantic delimiter
|
||||
// (preserving the sender's push boundary) and ECE state must be uniform
|
||||
// across a window (the same rule canAppend enforces for in-flow appends).
|
||||
// The IP-level ECN codepoint must also match: this check calls headersMatch
|
||||
// → ipHeadersMatch, which compares the full DSCP/ECN byte, so two slots with
|
||||
// differing ECN marks stay separate superpackets, each keeping its own mark.
|
||||
//
|
||||
// Note: a slot sealed by reorder (canAppend returned false on seq
|
||||
// mismatch) keeps psh=false, so this restriction does not block the
|
||||
// reorder-fix merge — only legitimate PSH-set seals.
|
||||
func canMergeSlots(prev, s *coalesceSlot) bool {
|
||||
if prev.psh {
|
||||
return false
|
||||
}
|
||||
if prev.fk != s.fk {
|
||||
return false
|
||||
}
|
||||
if prev.gsoSize != s.gsoSize {
|
||||
return false
|
||||
}
|
||||
if prev.nextSeq != slotSeedSeq(s) {
|
||||
return false
|
||||
}
|
||||
if prev.numSeg+s.numSeg > tcpCoalesceMaxSegs {
|
||||
return false
|
||||
}
|
||||
if prev.hdrLen+prev.totalPay+s.totalPay > tcpCoalesceBufSize {
|
||||
return false
|
||||
}
|
||||
if len(prev.payIovs[len(prev.payIovs)-1]) != prev.gsoSize {
|
||||
return false
|
||||
}
|
||||
prevFlags := prev.hdrBuf[prev.ipHdrLen+13]
|
||||
sFlags := s.hdrBuf[s.ipHdrLen+13]
|
||||
if (prevFlags^sFlags)&tcpFlagEce != 0 {
|
||||
return false
|
||||
}
|
||||
if !headersMatch(prev.hdrBuf[:prev.hdrLen], s.hdrBuf[:s.hdrLen], prev.isV6, prev.ipHdrLen) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// mergeSlots folds src into dst in place: payIovs concatenated, counters
|
||||
// and totals updated, PSH OR'd into the seed header so the push signal is
|
||||
// not lost. The seed header's seq, gsoSize, and fk are unchanged. Caller
|
||||
// is responsible for releasing src (it's no longer in c.slots after this call).
|
||||
func mergeSlots(dst, src *coalesceSlot) {
|
||||
dst.payIovs = append(dst.payIovs, src.payIovs...)
|
||||
dst.numSeg += src.numSeg
|
||||
dst.totalPay += src.totalPay
|
||||
dst.nextSeq = src.nextSeq
|
||||
if src.psh {
|
||||
dst.psh = true
|
||||
dst.hdrBuf[dst.ipHdrLen+13] |= tcpFlagPsh
|
||||
}
|
||||
}
|
||||
|
||||
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
|
||||
// already have its checksum field zeroed) and returns the folded/inverted
|
||||
// 16-bit value to store.
|
||||
@@ -717,9 +462,8 @@ func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
|
||||
return sum
|
||||
}
|
||||
|
||||
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it
|
||||
// unchanged (no one's complement). This is what virtio NEEDS_CSUM wants in
|
||||
// the L4 checksum field — the kernel will add the payload sum and invert.
|
||||
// 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)
|
||||
|
||||
@@ -2,9 +2,9 @@ package batch
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"runtime"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/test"
|
||||
)
|
||||
@@ -55,7 +55,31 @@ func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte {
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildICMPv4 returns a minimal non-TCP packet that takes the passthrough
|
||||
// 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)
|
||||
@@ -71,8 +95,7 @@ func buildICMPv4() []byte {
|
||||
// between batches, and reports per-packet cost.
|
||||
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
arena := NewArena(0)
|
||||
c := NewTCPCoalescer(nopTunWriter{}, test.NewLogger(), arena.Reserve, arena.Reset)
|
||||
c := newTestTCPCoalescer(b, nopTunWriter{})
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
@@ -113,8 +136,17 @@ func BenchmarkCommitInterleaved16(b *testing.B) {
|
||||
runCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseTCPBase
|
||||
// bails early and addPassthrough is the only work.
|
||||
// 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)
|
||||
@@ -126,7 +158,7 @@ func BenchmarkCommitPassthrough(b *testing.B) {
|
||||
|
||||
// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow.
|
||||
// Each packet takes the "TCP but not admissible" branch which does a
|
||||
// map delete + passthrough. Measures the seal-without-slot cost.
|
||||
// map delete + verbatim. Measures the seal-without-slot cost.
|
||||
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
||||
pay := make([]byte, 0)
|
||||
pkts := make([][]byte, 64)
|
||||
@@ -136,18 +168,24 @@ func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
||||
runCommitBench(b, pkts, 64)
|
||||
}
|
||||
|
||||
// runMultiCommitBench drives MultiCoalescer.Commit. The dispatcher does
|
||||
// the IP/L4 parse once and passes the parsed struct to the lane, so this
|
||||
// is the bench that shows the savings of skipping the lane's re-parse.
|
||||
// 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(), NewArena(0), true, true)
|
||||
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++ {
|
||||
pkt := pkts[i%len(pkts)]
|
||||
if err := m.Commit(pkt); err != nil {
|
||||
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 {
|
||||
@@ -174,68 +212,3 @@ func BenchmarkMultiCommitInterleaved4(b *testing.B) {
|
||||
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
||||
runMultiCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// flowKeyPair is one comparison input for the flowKeyCompare bench.
|
||||
type flowKeyPair struct{ a, b flowKey }
|
||||
|
||||
// makeFlowKey builds an IPv4 flowKey from compact inputs.
|
||||
func makeFlowKey(srcLow, dstLow uint32, sport, dport uint16) flowKey {
|
||||
var fk flowKey
|
||||
binary.BigEndian.PutUint32(fk.src[12:16], srcLow)
|
||||
binary.BigEndian.PutUint32(fk.dst[12:16], dstLow)
|
||||
fk.sport = sport
|
||||
fk.dport = dport
|
||||
return fk
|
||||
}
|
||||
|
||||
// flowKeyCases are the workload mixes flowKeyCompare sees in practice.
|
||||
// - sameFlow: equal keys; tests the equal-path cost (sort runs hit this
|
||||
// repeatedly when many segments share a flow).
|
||||
// - sportDiffers: same src/dst/dport, different sport — the typical
|
||||
// "sibling flows from one host to one server" pattern.
|
||||
// - dstDiffers: same src/sport/dport, different dst — outbound to many
|
||||
// servers from a fixed local port.
|
||||
// - allDiffer: every field differs; worst case for short-circuiting.
|
||||
func flowKeyCases() map[string][]flowKeyPair {
|
||||
const n = 64
|
||||
cases := map[string][]flowKeyPair{
|
||||
"sameFlow": make([]flowKeyPair, n),
|
||||
"sportDiffers": make([]flowKeyPair, n),
|
||||
"dstDiffers": make([]flowKeyPair, n),
|
||||
"allDiffer": make([]flowKeyPair, n),
|
||||
}
|
||||
for i := range n {
|
||||
base := makeFlowKey(0x0a000001, 0x0a000002, 40000, 443)
|
||||
cases["sameFlow"][i] = flowKeyPair{a: base, b: base}
|
||||
cases["sportDiffers"][i] = flowKeyPair{
|
||||
a: base,
|
||||
b: makeFlowKey(0x0a000001, 0x0a000002, uint16(40001+i), 443),
|
||||
}
|
||||
cases["dstDiffers"][i] = flowKeyPair{
|
||||
a: base,
|
||||
b: makeFlowKey(0x0a000001, uint32(0x0a000002+i+1), 40000, 443),
|
||||
}
|
||||
cases["allDiffer"][i] = flowKeyPair{
|
||||
a: makeFlowKey(uint32(0x0a000001+i), uint32(0x0a000002+i), uint16(40000+i), uint16(80+i)),
|
||||
b: makeFlowKey(uint32(0x0b000001+i), uint32(0x0b000002+i), uint16(50000+i), uint16(443+i)),
|
||||
}
|
||||
}
|
||||
return cases
|
||||
}
|
||||
|
||||
// BenchmarkFlowKeyCompare measures flowKeyCompare across the workloads
|
||||
// the sort step actually sees. Use this to compare reorderings.
|
||||
func BenchmarkFlowKeyCompare(b *testing.B) {
|
||||
for name, pairs := range flowKeyCases() {
|
||||
b.Run(name, func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
var sink int
|
||||
for i := 0; i < b.N; i++ {
|
||||
p := pairs[i&(len(pairs)-1)]
|
||||
sink += flowKeyCompare(p.a, p.b)
|
||||
}
|
||||
runtime.KeepAlive(sink)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+472
-330
File diff suppressed because it is too large
Load Diff
@@ -6,7 +6,7 @@ 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, outerECNs []byte) error
|
||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
|
||||
}
|
||||
|
||||
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
|
||||
@@ -16,7 +16,6 @@ type SendBatch struct {
|
||||
out batchWriter
|
||||
bufs [][]byte
|
||||
dsts []netip.AddrPort
|
||||
ecns []byte
|
||||
arena *Arena
|
||||
}
|
||||
|
||||
@@ -26,7 +25,6 @@ func NewSendBatch(out batchWriter, batchCap, arenaSize int) *SendBatch {
|
||||
out: out,
|
||||
bufs: make([][]byte, 0, batchCap),
|
||||
dsts: make([]netip.AddrPort, 0, batchCap),
|
||||
ecns: make([]byte, 0, batchCap),
|
||||
arena: NewArena(arenaSize),
|
||||
}
|
||||
}
|
||||
@@ -40,21 +38,22 @@ func (b *SendBatch) Reserve(sz int) []byte {
|
||||
// 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, outerECN byte) {
|
||||
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort) {
|
||||
b.bufs = append(b.bufs, pkt)
|
||||
b.dsts = append(b.dsts, dst)
|
||||
b.ecns = append(b.ecns, outerECN)
|
||||
}
|
||||
|
||||
func (b *SendBatch) Flush() error {
|
||||
// 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 {
|
||||
err = b.out.WriteBatch(b.bufs, b.dsts, b.ecns)
|
||||
written, err = b.out.WriteBatch(b.bufs, b.dsts)
|
||||
}
|
||||
clear(b.bufs)
|
||||
b.bufs = b.bufs[:0]
|
||||
b.dsts = b.dsts[:0]
|
||||
b.ecns = b.ecns[:0]
|
||||
b.arena.Reset()
|
||||
return err
|
||||
return written, err
|
||||
}
|
||||
|
||||
@@ -8,10 +8,9 @@ import (
|
||||
type fakeBatchWriter struct {
|
||||
bufs [][]byte
|
||||
addrs []netip.AddrPort
|
||||
ecns []byte
|
||||
}
|
||||
|
||||
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns []byte) error {
|
||||
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))
|
||||
@@ -21,8 +20,7 @@ func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns
|
||||
w.bufs[i] = cp
|
||||
}
|
||||
w.addrs = append(w.addrs[:0], addrs...)
|
||||
w.ecns = append(w.ecns[:0], ecns...)
|
||||
return nil
|
||||
return len(bufs), nil
|
||||
}
|
||||
|
||||
func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||
@@ -36,9 +34,9 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||
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, 0)
|
||||
b.Commit(pkt, ap)
|
||||
}
|
||||
if err := b.Flush(); err != nil {
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
if len(fw.bufs) != 4 {
|
||||
@@ -55,7 +53,7 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||
|
||||
// Flush again with nothing committed — should be a no-op.
|
||||
fw.bufs = nil
|
||||
if err := b.Flush(); err != nil {
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("empty Flush: %v", err)
|
||||
}
|
||||
if fw.bufs != nil {
|
||||
@@ -77,9 +75,9 @@ func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
||||
for i := 0; i < 3; i++ {
|
||||
s := b.Reserve(8)
|
||||
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
|
||||
b.Commit(pkt, ap, 0)
|
||||
b.Commit(pkt, ap)
|
||||
}
|
||||
if err := b.Flush(); err != nil {
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
|
||||
@@ -98,18 +96,18 @@ func TestSendBatchGrowPreservesCommitted(t *testing.T) {
|
||||
|
||||
s1 := b.Reserve(4)
|
||||
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
|
||||
b.Commit(pkt1, ap, 0)
|
||||
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, 0)
|
||||
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 {
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
if len(fw.bufs) != 2 {
|
||||
|
||||
+142
-152
@@ -1,6 +1,7 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
|
||||
@@ -18,72 +19,56 @@ const udpCoalesceBufSize = 65535
|
||||
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
|
||||
const udpCoalesceMaxSegs = 64
|
||||
|
||||
// udpCoalesceHdrCap is the scratch space we copy a seed's IP+UDP header
|
||||
// into. IPv6 (40) + UDP (8) = 48; round up for safety.
|
||||
const udpCoalesceHdrCap = 64
|
||||
|
||||
// udpSlot is one entry in the UDPCoalescer's ordered event queue.
|
||||
type udpSlot struct {
|
||||
passthrough bool
|
||||
rawPkt []byte // borrowed when passthrough
|
||||
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
|
||||
hdrBuf [udpCoalesceHdrCap]byte
|
||||
hdrLen int
|
||||
ipHdrLen int
|
||||
isV6 bool
|
||||
gsoSize int // per-segment UDP payload length
|
||||
numSeg int
|
||||
totalPay int
|
||||
// sealed closes the chain: set when a sub-gsoSize segment is appended
|
||||
// (kernel UDP-GSO requires every segment but the last to be exactly
|
||||
// gsoSize) or when limits are hit. No further appends after.
|
||||
sealed bool
|
||||
payIovs [][]byte
|
||||
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. Falls back to per-packet writes when the
|
||||
// underlying writer doesn't support USO.
|
||||
//
|
||||
// All output — coalesced or not — is deferred until Flush so per-flow
|
||||
// arrival order is preserved on the wire. Cross-flow order is NOT preserved
|
||||
// across the TCP/UDP/passthrough split when this coalescer runs alongside
|
||||
// others — see multi_coalesce.go. Per-flow order is preserved because a
|
||||
// single 5-tuple only ever lands in one lane and each lane preserves its
|
||||
// own slot order.
|
||||
// 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 {
|
||||
plainW io.Writer
|
||||
gsoW tio.GSOWriter // nil when the queue can't accept GSO_UDP_L4
|
||||
|
||||
w tio.GSOWriter
|
||||
slots []*udpSlot
|
||||
openSlots map[flowKey]*udpSlot
|
||||
pool []*udpSlot
|
||||
reserver Reserver
|
||||
resetter Resetter
|
||||
// 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
|
||||
}
|
||||
|
||||
// NewUDPCoalescer wraps w. The caller is responsible for only constructing
|
||||
// this when the underlying Queue's Capabilities advertise USO; otherwise
|
||||
// the kernel may reject GSO_UDP_L4 writes. If w does not implement
|
||||
// tio.GSOWriter at all (single-packet Queue), the coalescer degrades to
|
||||
// plain Writes — same defensive shape as the TCP coalescer.
|
||||
func NewUDPCoalescer(w io.Writer, reserver Reserver, resetter Resetter) *UDPCoalescer {
|
||||
c := &UDPCoalescer{
|
||||
plainW: w,
|
||||
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),
|
||||
reserver: reserver,
|
||||
resetter: resetter,
|
||||
}
|
||||
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP); ok {
|
||||
c.gsoW = gw
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// parsedUDP holds the fields extracted from a single parse so later steps
|
||||
@@ -95,104 +80,111 @@ type parsedUDP struct {
|
||||
payLen int
|
||||
}
|
||||
|
||||
// parseUDP extracts the flow key and IP/UDP offsets for a UDP packet.
|
||||
// Returns ok=false for non-UDP, malformed, or unsupported header shapes
|
||||
// (IPv4 with options/fragmentation, IPv6 with extension headers).
|
||||
func parseUDP(pkt []byte) (parsedUDP, bool) {
|
||||
var p parsedUDP
|
||||
ip, ok := parseIPPrologue(pkt, ipProtoUDP)
|
||||
// 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 p, false
|
||||
return false
|
||||
}
|
||||
pkt = ip.pkt
|
||||
p.fk = ip.fk
|
||||
p.ipHdrLen = ip.ipHdrLen
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
if len(pkt) < p.ipHdrLen+8 {
|
||||
return p, false
|
||||
// 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
|
||||
}
|
||||
p.hdrLen = p.ipHdrLen + 8
|
||||
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
|
||||
udpLen := int(binary.BigEndian.Uint16(pkt[p.ipHdrLen+4 : p.ipHdrLen+6]))
|
||||
if udpLen < 8 || udpLen > len(pkt)-p.ipHdrLen {
|
||||
return p, false
|
||||
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[p.ipHdrLen : p.ipHdrLen+2])
|
||||
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
||||
return p, true
|
||||
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
|
||||
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
|
||||
return true
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) Reserve(sz int) []byte {
|
||||
return c.reserver(sz)
|
||||
// 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)
|
||||
}
|
||||
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||
func (c *UDPCoalescer) Commit(pkt []byte) error {
|
||||
if c.gsoW == nil {
|
||||
c.addPassthrough(pkt)
|
||||
// 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
|
||||
}
|
||||
info, ok := parseUDP(pkt)
|
||||
if !ok {
|
||||
c.addPassthrough(pkt)
|
||||
var info parsedUDP
|
||||
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(pkt, info)
|
||||
return c.commitParsed(sp.pkt, &info)
|
||||
}
|
||||
|
||||
// commitParsed is the post-parse half of Commit. The caller must have
|
||||
// already verified parseUDP succeeded. Used by MultiCoalescer.Commit to
|
||||
// avoid re-walking the IP/UDP header.
|
||||
func (c *UDPCoalescer) commitParsed(pkt []byte, info parsedUDP) error {
|
||||
if c.gsoW == nil {
|
||||
c.addPassthrough(pkt)
|
||||
return nil
|
||||
}
|
||||
// A zero-length UDP datagram (UDP `length` == 8) is legal and must still
|
||||
// reach the TUN, but it can't be coalesced: a GSO slot would store an
|
||||
// empty payload iovec and the kernel has nothing to segment. Seal any
|
||||
// open chain for this flow (so a later, non-empty datagram seeds fresh
|
||||
// *after* this one and per-flow arrival order is preserved) and deliver
|
||||
// it as a plain single datagram.
|
||||
// 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 {
|
||||
delete(c.openSlots, info.fk)
|
||||
c.addPassthrough(pkt)
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
if open := c.openSlots[info.fk]; open != 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) {
|
||||
c.appendPayload(open, pkt, info)
|
||||
if open.sealed {
|
||||
delete(c.openSlots, info.fk)
|
||||
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 — seal it and fall through to seed a fresh slot.
|
||||
delete(c.openSlots, info.fk)
|
||||
// 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
|
||||
}
|
||||
|
||||
// Flush drains every queued slot and calls the configured Resetter.
|
||||
func (c *UDPCoalescer) Flush() error {
|
||||
first := c.drain()
|
||||
if c.resetter != nil {
|
||||
c.resetter()
|
||||
}
|
||||
return first
|
||||
}
|
||||
|
||||
// drain emits every queued slot in arrival order and clears the slot state.
|
||||
// It does NOT reset the arena: borrowed payload slices stay valid until the
|
||||
// arena's owner recycles it.
|
||||
func (c *UDPCoalescer) drain() error {
|
||||
var first error
|
||||
for _, s := range c.slots {
|
||||
var err error
|
||||
if s.passthrough {
|
||||
_, err = c.plainW.Write(s.rawPkt)
|
||||
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)
|
||||
}
|
||||
@@ -204,25 +196,38 @@ func (c *UDPCoalescer) drain() error {
|
||||
clear(c.slots)
|
||||
c.slots = c.slots[:0]
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
return first
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) addPassthrough(pkt []byte) {
|
||||
// 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.passthrough = true
|
||||
s.verbatim = true
|
||||
s.rawPkt = pkt
|
||||
c.slots = append(c.slots, s)
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) seed(pkt []byte, info parsedUDP) {
|
||||
if info.hdrLen > udpCoalesceHdrCap || info.hdrLen+info.payLen > udpCoalesceBufSize {
|
||||
c.addPassthrough(pkt)
|
||||
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.passthrough = false
|
||||
s.rawPkt = nil
|
||||
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
||||
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
|
||||
@@ -230,19 +235,16 @@ func (c *UDPCoalescer) seed(pkt []byte, info parsedUDP) {
|
||||
s.gsoSize = info.payLen
|
||||
s.numSeg = 1
|
||||
s.totalPay = info.payLen
|
||||
s.sealed = false
|
||||
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 s.sealed {
|
||||
return false
|
||||
}
|
||||
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info *parsedUDP) bool {
|
||||
if info.hdrLen != s.hdrLen {
|
||||
return false
|
||||
}
|
||||
@@ -255,20 +257,25 @@ func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info parsedUDP) bool {
|
||||
if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize {
|
||||
return false
|
||||
}
|
||||
if !udpHeadersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||
// 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
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info parsedUDP) {
|
||||
// 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
|
||||
if info.payLen < s.gsoSize {
|
||||
// Last-segment-can-be-shorter: this seals the chain.
|
||||
s.sealed = true
|
||||
}
|
||||
return info.payLen < s.gsoSize
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) take() *udpSlot {
|
||||
@@ -282,30 +289,19 @@ func (c *UDPCoalescer) take() *udpSlot {
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) release(s *udpSlot) {
|
||||
s.passthrough = false
|
||||
s.rawPkt = nil
|
||||
// Reset every field, identity ones included; see TCPCoalescer.release.
|
||||
clear(s.payIovs)
|
||||
s.payIovs = s.payIovs[:0]
|
||||
s.numSeg = 0
|
||||
s.totalPay = 0
|
||||
s.sealed = false
|
||||
*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 kernel's ip_rcv_core (v4) and
|
||||
// ip6_rcv_core (v6) trim the skb to those length fields, so per-segment
|
||||
// values would silently drop everything but the first segment. The kernel
|
||||
// then walks each segment in __udp_gso_segment, recomputing per-segment
|
||||
// uh->len / iph->tot_len / IPv6 plen and adjusting the checksum via
|
||||
// `check = csum16_add(csum16_sub(uh->check, uh->len), newlen)` — meaning
|
||||
// our seed's uh->check must be consistent with the seed's uh->len, which
|
||||
// is what passing the total to both pseudoSum and the UDP length field
|
||||
// guarantees.
|
||||
// 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.hdrBuf[:s.hdrLen]
|
||||
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)
|
||||
|
||||
@@ -330,14 +326,11 @@ func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
||||
udpCsumOff := s.ipHdrLen + 6
|
||||
binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum))
|
||||
|
||||
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP)
|
||||
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. Length
|
||||
// fields are masked out (flushSlot rewrites them), but the IP-level ECN
|
||||
// codepoint is compared (via ipHeadersMatch) so segments with differing ECN
|
||||
// don't coalesce, matching kernel GRO.
|
||||
// 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
|
||||
@@ -345,11 +338,8 @@ func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||
if !ipHeadersMatch(a, b, isV6) {
|
||||
return false
|
||||
}
|
||||
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8] —
|
||||
// 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
|
||||
if a[udp] != b[udp] || a[udp+1] != b[udp+1] || a[udp+2] != b[udp+2] || a[udp+3] != b[udp+3] {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
return bytes.Equal(a[udp:udp+4], b[udp:udp+4])
|
||||
}
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
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,7 +1,9 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"testing"
|
||||
)
|
||||
|
||||
@@ -58,29 +60,31 @@ func buildUDPv6(sport, dport uint16, payload []byte) []byte {
|
||||
return pkt
|
||||
}
|
||||
|
||||
func TestUDPCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: false}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 100))
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
// 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")
|
||||
}
|
||||
if len(w.writes) != 0 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("no Add-time writes: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
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 err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want single plain write, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
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}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
// ICMP packet
|
||||
pkt := make([]byte, 28)
|
||||
pkt[0] = 0x45
|
||||
@@ -101,8 +105,7 @@ func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
||||
|
||||
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 800))
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -110,17 +113,21 @@ func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Single-segment flush goes through WriteGSO; the writer infers GSO_NONE
|
||||
// from len(pays)==1 and the kernel fills in the UDP csum (NEEDS_CSUM).
|
||||
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
||||
// 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}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
for i := 0; i < 3; i++ {
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
@@ -160,8 +167,7 @@ func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
|
||||
// Last segment may be shorter, sealing the chain.
|
||||
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
full := make([]byte, 1200)
|
||||
tail := make([]byte, 600)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
@@ -180,22 +186,23 @@ func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (sealed + new seed), got %d", len(w.gsoWrites))
|
||||
// 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("first super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
||||
t.Errorf("super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
||||
}
|
||||
if len(w.gsoWrites[1].pays) != 1 {
|
||||
t.Errorf("second super: want 1 pay (re-seed), got %d", len(w.gsoWrites[1].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}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -205,16 +212,21 @@ func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 separate seeds, got %d", len(w.gsoWrites))
|
||||
// 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}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 800)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -245,8 +257,7 @@ func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
||||
// Caps at udpCoalesceMaxSegs.
|
||||
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
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 {
|
||||
@@ -271,12 +282,12 @@ func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
||||
|
||||
// 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 seeds a fresh superpacket that keeps CE; the
|
||||
// trailing Not-ECT datagram seeds another.
|
||||
// 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}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 800)
|
||||
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT)
|
||||
pkt1 := buildUDPv4(1000, 53, pay)
|
||||
@@ -290,16 +301,13 @@ func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 3 {
|
||||
t.Fatalf("want 3 separate seeds (differing ECN), got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||
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, g := range w.gsoWrites {
|
||||
if len(g.pays) != 1 {
|
||||
t.Errorf("gso %d pay count=%d want 1", i, len(g.pays))
|
||||
}
|
||||
if got := g.hdr[1] & 0x03; got != wantECN[i] {
|
||||
t.Errorf("gso %d ECN=%#x want %#x", i, got, wantECN[i])
|
||||
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])
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -307,8 +315,7 @@ func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
||||
// IPv6 path: same flow, equal-sized → coalesced.
|
||||
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
for i := 0; i < 3; i++ {
|
||||
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
|
||||
@@ -344,8 +351,7 @@ func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
||||
// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte.
|
||||
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 800)
|
||||
pkt0 := buildUDPv4(1000, 53, pay)
|
||||
pkt1 := buildUDPv4(1000, 53, pay)
|
||||
@@ -359,16 +365,16 @@ func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 separate seeds (different DSCP), got %d", len(w.gsoWrites))
|
||||
// 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}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
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 {
|
||||
@@ -389,8 +395,7 @@ func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
||||
// reach the GSO path. Regression: must not panic and must be written.
|
||||
func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -406,11 +411,10 @@ func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// IPv6 zero-length UDP datagram: same passthrough contract as v4.
|
||||
// IPv6 zero-length UDP datagram: same verbatim contract as v4.
|
||||
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -431,8 +435,7 @@ func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
|
||||
// wire — per-flow arrival order (full, empty, full) must be preserved.
|
||||
func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
full := make([]byte, 800)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -447,17 +450,23 @@ func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// The empty datagram sealed the first slot, so the trailing full packet
|
||||
// can't join it: two single-segment superpackets bracket one plain write.
|
||||
if len(w.gsoWrites) != 2 || len(w.writes) != 1 {
|
||||
t.Fatalf("want 2 gso writes + 1 plain, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||
// 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}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
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 {
|
||||
@@ -470,3 +479,58 @@ func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -8,37 +8,69 @@ import (
|
||||
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
)
|
||||
|
||||
// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
|
||||
// seeds and a handful of starting alignments, asserting that our local
|
||||
// Checksum matches gvisor's reference bit-for-bit.
|
||||
func TestChecksumMatchesGvisor(t *testing.T) {
|
||||
rng := rand.New(rand.NewPCG(1, 2))
|
||||
const padFront = 16
|
||||
// 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
|
||||
}
|
||||
|
||||
// Random pool large enough for the longest case + alignment slop.
|
||||
pool := make([]byte, 4096+padFront)
|
||||
for i := range pool {
|
||||
pool[i] = byte(rng.Uint32())
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
|
||||
seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc}
|
||||
offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16}
|
||||
// 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
|
||||
|
||||
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 := Checksum(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
|
||||
length, off, seed, got, want)
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -46,23 +78,28 @@ func TestChecksumMatchesGvisor(t *testing.T) {
|
||||
// historically tripped up checksum implementations: all-zero, all-0xff,
|
||||
// alternating, and ascending sequences.
|
||||
func TestChecksumPatternedBuffers(t *testing.T) {
|
||||
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 := Checksum(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
|
||||
name, length, seed, got, want)
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -98,36 +135,41 @@ func ascending(n int) []byte {
|
||||
// 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) {
|
||||
rng := rand.New(rand.NewPCG(42, 17))
|
||||
const padFront = 16
|
||||
const maxK = 8
|
||||
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())
|
||||
}
|
||||
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
|
||||
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 := Checksum(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)
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
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},
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
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},
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
//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
|
||||
@@ -9,10 +9,9 @@ import (
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// blockOn parks the calling goroutine until fd is ready (events is POLLIN for
|
||||
// reads, POLLOUT for writes) or shutdownFd signals teardown. It builds the
|
||||
// pollfd array on the stack every call, so concurrent callers on the same
|
||||
// Queue never share Revents storage.
|
||||
// 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).
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
@@ -17,18 +18,17 @@ type offloadQueueSet struct {
|
||||
// pqi is exactly the same as pq, but stored as the interface type
|
||||
pqi []Queue
|
||||
shutdownFd int
|
||||
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6
|
||||
// with the kernel. Queues created by Add inherit this and surface it
|
||||
// via Offload.USOSupported so coalescers can gate USO emission.
|
||||
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6 with the kernel.
|
||||
// Queues created by Add inherit this and surface it via Offload.USOSupported so coalescers can gate USO emission.
|
||||
usoEnabled bool
|
||||
closed atomic.Bool
|
||||
// l is handed to each queue for its bad-vnet-header drop logging.
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do
|
||||
// TSO segmentation in userspace. usoEnabled tells downstream queues whether
|
||||
// the kernel agreed to deliver/accept GSO_UDP_L4 superpackets — coalescers
|
||||
// should fall back to per-packet writes when this is false.
|
||||
func NewOffloadQueueSet(usoEnabled bool) (QueueSet, error) {
|
||||
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do TSO segmentation.
|
||||
// usoEnabled tells downstream queues whether the kernel agreed to deliver/accept GSO_UDP_L4 superpackets.
|
||||
func NewOffloadQueueSet(usoEnabled bool, l *slog.Logger) (QueueSet, error) {
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||
@@ -39,6 +39,7 @@ func NewOffloadQueueSet(usoEnabled bool) (QueueSet, error) {
|
||||
pqi: []Queue{},
|
||||
shutdownFd: shutdownFd,
|
||||
usoEnabled: usoEnabled,
|
||||
l: l,
|
||||
}
|
||||
|
||||
return out, nil
|
||||
@@ -49,7 +50,10 @@ func (c *offloadQueueSet) Queues() []Queue {
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) Add(fd int) error {
|
||||
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled)
|
||||
if c.closed.Load() {
|
||||
return errors.New("queue set already closed")
|
||||
}
|
||||
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled, c.l)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -73,23 +77,21 @@ func (c *offloadQueueSet) Close() error {
|
||||
|
||||
errs := []error{}
|
||||
|
||||
// Signal all readers blocked in poll to wake up and exit. They observe
|
||||
// POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||
// Signal all readers blocked in poll to wake up and exit.
|
||||
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||
if err := c.wakeForShutdown(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
||||
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
|
||||
// to this container.
|
||||
for _, x := range c.pq {
|
||||
if err := x.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references
|
||||
// it, so it must outlive the wake + per-queue teardown above.
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references it,
|
||||
// so it must outlive the wake + per-queue teardown above.
|
||||
if err := unix.Close(c.shutdownFd); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user