Merge main into poc/certificate-posture

This commit is contained in:
Viktor Liu
2026-10-05 19:09:43 +02:00
390 changed files with 28443 additions and 14545 deletions
+22
View File
@@ -46,3 +46,25 @@ updates:
wireguard: wireguard:
patterns: patterns:
- "golang.zx2c4.com/wireguard*" - "golang.zx2c4.com/wireguard*"
# Base images of the source-build Dockerfiles, pinned by digest (Chainguard
# publishes only :latest for free). Dockerfile.release files feed goreleaser
# and keep the published images as they are, so their bases are left alone.
- package-ecosystem: "docker"
directories:
- "/upload-server"
schedule:
interval: "weekly"
open-pull-requests-limit: 3
groups:
base-images:
patterns:
- "*"
ignore:
- dependency-name: "gcr.io/distroless/base"
# Go minor and major versions move with the rest of the repository;
# patch releases and new digests of the pinned tag still come through.
- dependency-name: "golang"
update-types:
- "version-update:semver-minor"
- "version-update:semver-major"
+338
View File
@@ -0,0 +1,338 @@
#!/usr/bin/env bash
set -euo pipefail
fail() {
echo "::error::$*" >&2
exit 1
}
if [[ ${RUNNER_ENVIRONMENT:-} != github-hosted || ${RUNNER_OS:-} != macOS || $(uname -s) != Darwin ]]; then
fail "This test installs a system daemon and must run on a disposable GitHub macOS runner."
fi
if [[ $EUID == 0 ]]; then
fail "Run this script as the Homebrew user, not root."
fi
readonly test_dir="${RUNNER_TEMP:?}/homebrew-cask"
readonly results_dir="$test_dir/results"
readonly app='/Applications/Netbird UI.app'
readonly plist='/Library/LaunchDaemons/netbird.plist'
readonly cask='netbirdio/tap/netbird-ui'
readonly formula='netbirdio/tap/netbird'
readonly published_cask="$test_dir/published-netbird-ui.rb"
readonly legacy_cask="$test_dir/legacy-netbird-ui.rb"
readonly rendered_cask="$test_dir/rendered-netbird-ui.rb"
readonly fixture_dir="$test_dir/fixture"
readonly serve_dir="$test_dir/serve"
readonly fixture_zip="$serve_dir/netbird-ui.zip"
readonly fixture_port=18080
readonly fixture_url="http://127.0.0.1:$fixture_port/netbird-ui.zip"
readonly marker="$test_dir/installer.marker"
mkdir -p "$results_dir" "$fixture_dir/netbird_ui_darwin" "$serve_dir" "$test_dir/downloads"
exec > >(tee "$results_dir/test.log") 2>&1
sudo -n true
if command -v netbird || [[ -e "$app" || -e "$plist" ]] || pgrep -x netbird-ui; then
fail "The runner already has NetBird installed or running."
fi
if sudo launchctl print system/netbird > "$results_dir/initial-service.log" 2>&1; then
fail "The runner already has a NetBird service loaded."
fi
install_attempted=false
server_pid=''
daemon_pid=''
version=''
stop_ui() {
local status=0
sudo pkill -x netbird-ui || status=$?
# pkill returns 1 when the UI is already closed.
[[ $status == 0 || $status == 1 ]]
}
cleanup() {
local status=$?
trap - EXIT
set +e
if [[ $install_attempted == true ]]; then
stop_ui || status=1
if [[ -S /var/run/netbird.sock ]]; then
sudo netbird down || status=1
fi
if brew list --cask "$cask" >/dev/null 2>&1 || [[ -e "$app" ]]; then
brew uninstall --cask --force "$cask" || status=1
fi
# A failed cask install can leave a daemon even after Homebrew rolls back the app.
if sudo launchctl print system/netbird > "$results_dir/cleanup-service.log" 2>&1; then
sudo netbird service stop || status=1
fi
if [[ -e "$plist" ]]; then
sudo netbird service uninstall || status=1
fi
fi
if [[ -f /var/log/netbird/client.log ]]; then
sudo cat /var/log/netbird/client.log > "$results_dir/client.log" || status=1
fi
if command -v netbird >/dev/null; then
brew uninstall --formula "$formula" || status=1
fi
if [[ -n $server_pid ]]; then
kill "$server_pid" 2>/dev/null || true
fi
exit "$status"
}
trap cleanup EXIT
trap 'exit 130' INT
trap 'exit 143' TERM
run_logged() {
local name=$1
shift
"$@" 2>&1 | tee "$results_dir/$name.log"
}
cask_field() {
local stanza=$1 file=$2
sed -nE "s/^[[:space:]]*$stanza \"([^\"]+)\".*/\\1/p" "$file"
}
release_fields() {
local file=$1
grep -E '^[[:space:]]*(version|url|sha256|app) ' "$file"
}
use_cask() {
local file=$1
cp "$file" "$tap_dir/Casks/netbird-ui.rb"
}
# The released installer opens the UI as root, which never returns on a headless
# runner. The cask only needs two script paths and a version argument, so the test
# ships a stub bundle that records what it received and starts the daemon.
build_fixture() {
local bundle="$fixture_dir/netbird_ui_darwin"
printf '#!/bin/sh\nexit 0\n' > "$bundle/netbird-ui"
chmod 755 "$bundle/netbird-ui"
# After a bootout launchd keeps tearing the previous daemon down for a couple of
# seconds, and loading the same label again fails until that finishes.
cat > "$bundle/installer.sh" <<EOF
#!/bin/sh
set -eu
export PATH=\$PATH:/usr/local/bin:/opt/homebrew/bin
printf 'version=%s\\nuid=%s\\n' "\$1" "\$(id -u)" > '$marker'
netbird service install
attempt=0
until netbird service start; do
attempt=\$((attempt + 1))
[ "\$attempt" -lt 15 ] || exit 1
sleep 1
done
EOF
printf '#!/bin/sh\nexit 0\n' > "$bundle/uninstaller.sh"
# Shipped without the executable bit so the 0755 seen after install can only come from the cask.
chmod 644 "$bundle/installer.sh" "$bundle/uninstaller.sh"
rm -f "$fixture_zip"
(cd "$fixture_dir" && zip -qr "$fixture_zip" netbird_ui_darwin)
}
start_fixture_server() {
python3 -m http.server "$fixture_port" --bind 127.0.0.1 --directory "$serve_dir" \
> "$results_dir/fixture-server.log" 2>&1 &
server_pid=$!
local attempt
for attempt in {1..20}; do
if curl --silent --fail --output /dev/null "$fixture_url"; then
return
fi
sleep 0.5
done
fail "The fixture HTTP server did not come up on port $fixture_port."
}
assert_published_layout() {
local url archive script
while read -r url; do
archive="$test_dir/downloads/${url##*/}"
curl --fail --location --silent --retry 3 --output "$archive" "$url"
for script in installer.sh uninstaller.sh; do
unzip -l "$archive" | grep -q " netbird_ui_darwin/$script\$" ||
fail "The published archive ${url##*/} has no netbird_ui_darwin/$script."
done
done < <(cask_field url "$published_cask")
}
assert_no_deprecations() {
if grep -Ei '(postflight|uninstall_preflight).*deprecated|deprecated.*(postflight|uninstall_preflight)' "$@"; then
fail "Homebrew reported a deprecated cask lifecycle hook."
fi
}
wait_for_daemon() {
local attempt
for attempt in {1..30}; do
if sudo launchctl print system/netbird > "$results_dir/service.log" 2>&1 &&
grep -Eq '^[[:space:]]*state = running$' "$results_dir/service.log"; then
return
fi
sleep 1
done
cat "$results_dir/service.log"
fail "The installed daemon did not reach the running state."
}
wait_for_exit() {
local pid=$1 attempt
for attempt in {1..30}; do
if ! sudo kill -0 "$pid" 2>/dev/null; then
return
fi
sleep 1
done
fail "Daemon process $pid is still running after removal."
}
assert_service_absent() {
if sudo launchctl print system/netbird > "$results_dir/removed-service.log" 2>&1; then
fail "The NetBird service is still loaded after removal."
fi
}
assert_installed() {
local script
[[ -f $marker ]] || fail "The cask did not run installer.sh."
grep -qx "version=$version" "$marker" || fail "installer.sh did not receive the cask version: $(cat "$marker")"
grep -qx 'uid=0' "$marker" || fail "installer.sh did not run as root: $(cat "$marker")"
[[ -d "$app" && -x "$app/netbird-ui" ]] || fail "The UI was not installed."
for script in installer.sh uninstaller.sh; do
[[ $(stat -f '%Lp' "$app/$script") == 755 ]] || fail "Incorrect permissions on $script."
done
[[ -f "$plist" ]] || fail "The installer did not create the daemon plist."
wait_for_daemon
daemon_pid=$(awk '/^[[:space:]]*pid = / { print $3; exit }' "$results_dir/service.log")
[[ $daemon_pid =~ ^[0-9]+$ ]] || fail "The running daemon has no PID."
sudo kill -0 "$daemon_pid"
}
assert_uninstalled() {
local log=$1
assert_no_deprecations "$log"
[[ ! -e "$app" ]] || fail "The UI app remains after uninstall."
[[ ! -e "$plist" ]] || fail "The daemon plist remains after uninstall."
assert_service_absent
wait_for_exit "$daemon_pid"
[[ $(netbird version) == "$version" ]] || fail "Cask uninstall removed the CLI dependency."
}
installed_caskfiles() {
local extension=$1
find "$(brew --caskroom)/netbird-ui/.metadata" -name "netbird-ui.$extension" 2>/dev/null
}
assert_legacy_metadata() {
installed_caskfiles rb | grep -q . || fail "The legacy cask did not leave a Ruby caskfile behind."
}
assert_steps_metadata() {
if installed_caskfiles rb | grep -q .; then
fail "Homebrew still keeps the legacy Ruby caskfile after reinstall."
fi
installed_caskfiles json | grep -q . || fail "Homebrew did not save the reinstalled cask as JSON."
}
brew --version
sw_vers
brew tap netbirdio/tap "${GITHUB_WORKSPACE:?}/.homebrew-cask-tap"
tap_dir=$(brew --repository netbirdio/tap)
readonly tap_dir
[[ -f "$tap_dir/Casks/netbird-ui.rb" ]] || fail "The tap has no Casks/netbird-ui.rb."
cp "$tap_dir/Casks/netbird-ui.rb" "$published_cask"
cp "$published_cask" "$results_dir/published-netbird-ui.rb"
version=$(brew info --json=v2 --formula "$formula" | jq -r '.formulae[0].versions.stable')
readonly version
[[ -n $version && $version != null ]] || fail "Could not read the formula version from the tap."
assert_published_layout
build_fixture
fixture_sha=$(shasum -a 256 "$fixture_zip" | cut -d' ' -f1)
readonly fixture_sha
start_fixture_server
export PROJECT=netbird-ui VERSION="$version"
export AMD="$fixture_zip" ARM="$fixture_zip" AMD_URL="$fixture_url" ARM_URL="$fixture_url"
gomplate -f "$GITHUB_WORKSPACE/client/ui/netbird-ui.rb.tmpl" -o "$rendered_cask"
cp "$rendered_cask" "$results_dir/rendered-netbird-ui.rb"
sed -E "s|^([[:space:]]*version) \"[^\"]+\"|\\1 \"$version\"|; s|^([[:space:]]*url) \"[^\"]+\"|\\1 \"$fixture_url\"|; s|^([[:space:]]*sha256) \"[^\"]+\"|\\1 \"$fixture_sha\"|" \
"$published_cask" > "$legacy_cask"
cp "$legacy_cask" "$results_dir/legacy-netbird-ui.rb"
if ! diff <(release_fields "$legacy_cask") <(release_fields "$rendered_cask"); then
fail "The rendered cask changes release data, not only lifecycle stanzas."
fi
use_cask "$rendered_cask"
brew info --json=v2 --cask "$cask" > "$results_dir/cask.json" 2> "$results_dir/load.log"
cat "$results_dir/load.log"
assert_no_deprecations "$results_dir/load.log"
run_logged style brew style --cask --only-cops=Cask/InstallSteps "$cask"
run_logged install-cli brew install --formula "$formula"
[[ $(netbird version) == "$version" ]] || fail "The installed CLI does not report the formula version."
for scenario in running stopped missing; do
echo "::group::Uninstall with $scenario service"
install_attempted=true
sudo rm -f "$marker"
run_logged "install-$scenario" brew install --cask "$cask"
assert_no_deprecations "$results_dir/install-$scenario.log"
assert_installed
stop_ui
case "$scenario" in
running) ;;
stopped)
run_logged stop-daemon sudo netbird service stop
wait_for_exit "$daemon_pid"
[[ -f "$plist" ]] || fail "Stopping the daemon unexpectedly removed its plist."
;;
missing)
run_logged stop-missing-daemon sudo netbird service stop
run_logged remove-daemon sudo netbird service uninstall
wait_for_exit "$daemon_pid"
[[ ! -e "$plist" ]] || fail "The missing-service scenario still has a plist."
assert_service_absent
;;
*) fail "Unknown uninstall scenario: $scenario" ;;
esac
run_logged "uninstall-$scenario" brew uninstall --cask "$cask"
assert_uninstalled "$results_dir/uninstall-$scenario.log"
echo "::endgroup::"
done
# Every existing user first meets the new cask through an upgrade of the published
# one, whose legacy flight blocks Homebrew replays from the saved Ruby caskfile.
echo "::group::Reinstall over the published legacy cask"
install_attempted=true
use_cask "$legacy_cask"
sudo rm -f "$marker"
run_logged install-legacy brew install --cask "$cask"
assert_installed
assert_legacy_metadata
stop_ui
use_cask "$rendered_cask"
sudo rm -f "$marker"
run_logged reinstall-legacy brew reinstall --cask "$cask"
assert_installed
assert_steps_metadata
stop_ui
run_logged uninstall-legacy brew uninstall --cask "$cask"
assert_uninstalled "$results_dir/uninstall-legacy.log"
echo "::endgroup::"
+3 -3
View File
@@ -38,12 +38,12 @@ jobs:
persist-credentials: false persist-credentials: false
- name: Set up Node.js - name: Set up Node.js
uses: actions/setup-node@v4 uses: actions/setup-node@v7
with: with:
node-version: "22" node-version: "22"
- name: Set up pnpm - name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0 uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with: with:
version: 11 version: 11
@@ -79,7 +79,7 @@ jobs:
run: echo "path=$(pnpm store path --silent)" >> "$GITHUB_OUTPUT" run: echo "path=$(pnpm store path --silent)" >> "$GITHUB_OUTPUT"
- name: Cache pnpm store - name: Cache pnpm store
uses: actions/cache@v4 uses: actions/cache@v6
with: with:
path: ${{ steps.pnpm-store.outputs.path }} path: ${{ steps.pnpm-store.outputs.path }}
key: ${{ runner.os }}-pnpm-${{ hashFiles('client/ui/frontend/pnpm-lock.yaml') }} key: ${{ runner.os }}-pnpm-${{ hashFiles('client/ui/frontend/pnpm-lock.yaml') }}
+1 -1
View File
@@ -33,7 +33,7 @@ jobs:
with: with:
usesh: true usesh: true
copyback: false copyback: false
release: "15.0" release: "15.1"
envs: "GO_VERSION" envs: "GO_VERSION"
prepare: | prepare: |
pkg install -y curl pkgconf xorg pkg install -y curl pkgconf xorg
+46
View File
@@ -80,3 +80,49 @@ jobs:
skip-save-cache: true skip-save-cache: true
cache-invalidation-interval: 0 cache-invalidation-interval: 0
args: --timeout=20m args: --timeout=20m
# Separate job rather than extra rows in the matrix above: those rows pick a
# GOOS by picking a runner OS, while android/ios are cross-compiled from
# ubuntu — an `include` entry with os: ubuntu-latest would merge into the
# Linux row instead of adding one. The package path is restricted because a
# whole-repo run under GOOS=android pulls *_linux.go files into packages that
# have no android counterpart.
golangci-mobile:
strategy:
fail-fast: false
matrix:
include:
- goos: android
goarch: arm64
packages: ./client/android/...
display_name: Android
- goos: ios
goarch: arm64
packages: ./client/ios/...
display_name: iOS
name: ${{ matrix.display_name }}
runs-on: ubuntu-latest
timeout-minutes: 25
env:
CGO_ENABLED: 0
GOOS: ${{ matrix.goos }}
GOARCH: ${{ matrix.goarch }}
steps:
- name: Checkout code
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
with:
persist-credentials: false
- name: Install Go
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
with:
go-version-file: "go.mod"
cache: false
- name: golangci-lint
uses: golangci/golangci-lint-action@82606bf257cbaff209d206a39f5134f0cfbfd2ee #v9.2.1
with:
version: latest
install-mode: binary
skip-cache: true
skip-save-cache: true
cache-invalidation-interval: 0
args: --timeout=20m ${{ matrix.packages }}
@@ -0,0 +1,64 @@
name: Mobile
on:
push:
branches:
- main
- "release-*"
pull_request:
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }}
cancel-in-progress: true
jobs:
android_build:
name: "Android / Build"
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
goarch: [arm64, arm, amd64, "386"]
env:
CGO_ENABLED: 0
GOOS: android
GOARCH: ${{ matrix.goarch }}
steps:
- name: Checkout repository
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
with:
persist-credentials: false
- name: Install Go
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
with:
go-version-file: "go.mod"
- name: Build Android bridge
run: go build ./client/android/...
- name: Vet Android bridge
if: matrix.goarch == 'arm64'
run: go vet ./client/android/...
ios_build:
name: "iOS / Build"
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
goarch: [arm64, amd64]
env:
CGO_ENABLED: 0
GOOS: ios
GOARCH: ${{ matrix.goarch }}
steps:
- name: Checkout repository
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
with:
persist-credentials: false
- name: Install Go
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
with:
go-version-file: "go.mod"
# No `go vet` counterpart: every ios target requires external (cgo)
# linking, which needs an Xcode toolchain the runner does not have.
- name: Build iOS SDK
run: go build ./client/ios/...
+199
View File
@@ -0,0 +1,199 @@
name: Red Hat Certification
# Certify published UBI images in the Red Hat Ecosystem Catalog. Called by
# release.yml on stable tags, or run by hand to (re)certify any released
# version. preflight submits every architecture of an image's manifest list
# to Pyxis; auto-publish on the component makes it public once certified.
#
# Each component's Partner Connect ID comes from the REDHAT_CERT_ID_<NAME>
# repository variable, e.g. REDHAT_CERT_ID_CLIENT_ROOTLESS. The run fails
# before certifying anything if a selected component's variable is not set.
on:
workflow_call:
inputs:
component:
type: string
required: true
version:
type: string
required: true
secrets:
PYXIS_API_TOKEN:
required: true
workflow_dispatch:
inputs:
component:
description: "Component to certify"
type: choice
required: true
default: all
options:
- all
- client-rootless
- reverse-proxy
version:
description: "Released version, e.g. v0.80.0"
type: string
required: true
permissions:
contents: read
jobs:
resolve:
name: Resolve components
runs-on: ubuntu-24.04
outputs:
version: ${{ steps.resolve.outputs.version }}
matrix: ${{ steps.resolve.outputs.matrix }}
steps:
- name: Resolve components and images
id: resolve
env:
COMPONENT: ${{ inputs.component }}
INPUT_VERSION: ${{ inputs.version }}
REPO_VARS: ${{ toJSON(vars) }}
run: |
set -euo pipefail
version="${INPUT_VERSION#v}"
if [[ ! "$version" =~ ^[0-9]+\.[0-9]+\.[0-9]+$ ]]; then
echo "::error::Only stable x.y.z versions are certified, got '${INPUT_VERSION}'"
exit 1
fi
# name, image repository, tag suffix (must match .goreleaser.yaml).
# Keep the names in sync with the workflow_dispatch options above.
components=(
"client-rootless ghcr.io/netbirdio/netbird -rootless-ubi"
"reverse-proxy ghcr.io/netbirdio/reverse-proxy -ubi"
)
matrix="[]"
missing=()
for c in "${components[@]}"; do
read -r name repo suffix <<< "$c"
[[ "$COMPONENT" == all || "$COMPONENT" == "$name" ]] || continue
var="REDHAT_CERT_ID_${name^^}"; var="${var//-/_}"
id="$(jq -r --arg v "$var" '.[$v] // empty' <<< "$REPO_VARS")"
if [[ -z "$id" ]]; then
missing+=("$var")
continue
fi
matrix="$(jq -c --arg n "$name" --arg t "${version}${suffix}" --arg r "${repo}:${version}${suffix}" --arg i "$id" \
'. + [{component: $n, tag: $t, ref: $r, component_id: $i}]' <<< "$matrix")"
done
if (( ${#missing[@]} )); then
echo "::error::Set these repository variables to the Partner Connect component IDs: ${missing[*]}"
exit 1
fi
if [[ "$matrix" == "[]" ]]; then
echo "::error::No component to certify for '${COMPONENT}'"
exit 1
fi
echo "Components to certify: ${matrix}"
echo "version=${version}" >> "$GITHUB_OUTPUT"
echo "matrix=${matrix}" >> "$GITHUB_OUTPUT"
certify:
name: "Certify ${{ matrix.component }} UBI image"
needs: resolve
runs-on: ubuntu-24.04
strategy:
fail-fast: false
matrix:
include: ${{ fromJSON(needs.resolve.outputs.matrix) }}
env:
PREFLIGHT_VERSION: "1.21.0"
# sha256 of preflight-linux-amd64 from the 1.21.0 GitHub release.
# Red Hat publishes no checksum file, so the value is pinned here.
PREFLIGHT_SHA256: "5e653135503c72f8702bbe31d7643197d12937c68086879133dd6b9650a9a449"
steps:
- name: Verify the multi-arch image is on ghcr.io
env:
IMAGE_REF: ${{ matrix.ref }}
run: |
set -euo pipefail
docker buildx imagetools inspect "$IMAGE_REF" --raw > manifest.json
for arch in amd64 arm64; do
if ! jq -e --arg a "$arch" '.manifests[] | select(.platform.architecture == $a)' manifest.json > /dev/null; then
echo "::error::${IMAGE_REF} has no ${arch} manifest"
exit 1
fi
done
echo "Manifest list for ${IMAGE_REF}:"
jq -r '.manifests[] | "\(.platform.os)/\(.platform.architecture) \(.digest)"' manifest.json
- name: Install preflight
run: |
set -euo pipefail
curl -fsSL --proto '=https' --proto-redir '=https' -o preflight \
"https://github.com/redhat-openshift-ecosystem/openshift-preflight/releases/download/${PREFLIGHT_VERSION}/preflight-linux-amd64"
echo "${PREFLIGHT_SHA256} preflight" | sha256sum -c -
chmod +x preflight
./preflight --version
- name: Run preflight checks and submit to Red Hat
env:
IMAGE_REF: ${{ matrix.ref }}
PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
PFLT_CERTIFICATION_COMPONENT_ID: ${{ matrix.component_id }}
PFLT_ARTIFACTS: artifacts
PFLT_LOGFILE: artifacts/preflight.log
PFLT_LOGLEVEL: info
PFLT_JUNIT: "true"
run: |
set -euo pipefail
# No --platform: preflight walks the manifest list and submits every
# architecture in one run, grouped under one manifest-list digest.
# preflight does not create the PFLT_LOGFILE directory, and --submit
# fails if the log file is missing.
mkdir -p artifacts
./preflight check container "$IMAGE_REF" --submit
- name: Fail if any check did not pass
run: |
set -euo pipefail
shopt -s nullglob
results=(artifacts/results.json artifacts/*/results.json)
if [[ ${#results[@]} -eq 0 ]]; then
echo "::error::preflight produced no results.json"
exit 1
fi
status=0
for f in "${results[@]}"; do
arch="$(basename "$(dirname "$f")")"
passed="$(jq -r '.passed' "$f")"
failed="$(jq -r '[.results.failed[]?.name] | join(", ")' "$f")"
echo "${arch}: passed=${passed} ${failed:+failed checks: ${failed}}"
[[ "$passed" == "true" ]] || status=1
done
exit $status
- name: Upload preflight artifacts
if: always()
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
with:
name: redhat-preflight-${{ matrix.component }}-${{ needs.resolve.outputs.version }}
path: artifacts/
retention-days: 30
- name: Wait for Pyxis to mark both architectures certified
env:
TAG: ${{ matrix.tag }}
COMPONENT_ID: ${{ matrix.component_id }}
PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
run: |
set -euo pipefail
# Filter on the tag server-side so older versions are found past the first page.
url="https://catalog.redhat.com/api/containers/v1/projects/certification/id/${COMPONENT_ID}/images?filter=repositories.tags.name==${TAG}&page_size=100"
for attempt in $(seq 1 20); do
certified="$(curl -fsS --proto '=https' --proto-redir '=https' -H "X-API-KEY: ${PFLT_PYXIS_API_TOKEN}" "$url" \
| jq -r --arg t "$TAG" '[.data[] | select(.repositories[]?.tags[]?.name == $t) | select(.certified == true) | .architecture] | unique | join(",")')"
echo "attempt ${attempt}: certified architectures for ${TAG}: ${certified:-none}"
if [[ "$certified" == "amd64,arm64" ]]; then
echo "Both architectures certified. Auto-publish is enabled on the component, so the catalog updates on its own."
exit 0
fi
sleep 30
done
echo "::error::Pyxis has not marked both architectures certified after 10 minutes. Check https://connect.redhat.com/component/view/${COMPONENT_ID}/images"
exit 1
+46 -11
View File
@@ -69,7 +69,7 @@ jobs:
with: with:
usesh: true usesh: true
copyback: false copyback: false
release: "15.0" release: "15.1"
envs: "GO_VERSION" envs: "GO_VERSION"
prepare: | prepare: |
# Install required packages # Install required packages
@@ -191,6 +191,17 @@ jobs:
# requires a changelog. Generated, not committed (see .gitignore). # requires a changelog. Generated, not committed (see .gitignore).
# chglog is a go.mod tool directive, so go.sum pins it and its deps. # chglog is a go.mod tool directive, so go.sum pins it and its deps.
run: bash release_files/rpm-changelog.sh run: bash release_files/rpm-changelog.sh
- name: Fill the RPM ISA provide version
# nfpm cannot emit rpmbuild's ISA provide and GoReleaser cannot template it.
run: bash release_files/rpm-provides.sh
- name: Set up Node.js
uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7.0.0
with:
node-version: '22'
- name: Install proxy web dependencies for license collection
# proxy/collect-licenses.sh reads the UI's license terms from node_modules.
working-directory: proxy/web
run: npm ci --ignore-scripts
- name: Set up QEMU - name: Set up QEMU
uses: docker/setup-qemu-action@06116385d9baf250c9f4dcb4858b16962ea869c3 #v4.1.0 uses: docker/setup-qemu-action@06116385d9baf250c9f4dcb4858b16962ea869c3 #v4.1.0
- name: Set up Docker Buildx - name: Set up Docker Buildx
@@ -230,14 +241,18 @@ jobs:
uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2 uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2
with: with:
version: ${{ env.GORELEASER_VER }} version: ${{ env.GORELEASER_VER }}
args: release --clean ${{ env.flags }} args: release --config .goreleaser.generated.yaml --clean ${{ env.flags }}
env: env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
HOMEBREW_TAP_GITHUB_TOKEN: ${{ secrets.HOMEBREW_TAP_GITHUB_TOKEN }} HOMEBREW_TAP_GITHUB_TOKEN: ${{ secrets.HOMEBREW_TAP_GITHUB_TOKEN }}
UPLOAD_DEBIAN_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }} UPLOAD_DEBIAN_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }}
UPLOAD_YUM_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }} UPLOAD_YUM_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }}
GPG_RPM_KEY_FILE: ${{ env.GPG_RPM_KEY_FILE }} GPG_RPM_KEY_FILE: ${{ env.GPG_RPM_KEY_FILE }}
NFPM_NETBIRD_RPM_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }} # One per nfpm id: GoReleaser looks the passphrase up as NFPM_<ID>_PASSPHRASE.
NFPM_NETBIRD_RPM_AMD64_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
NFPM_NETBIRD_RPM_ARM64_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
NFPM_NETBIRD_RPM_ARM_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
NFPM_NETBIRD_RPM_386_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
SKIP_PUBLISH: ${{ env.SKIP_PUBLISH }} SKIP_PUBLISH: ${{ env.SKIP_PUBLISH }}
SKIP_DOCKER_PUSH: ${{ env.SKIP_DOCKER_PUSH }} SKIP_DOCKER_PUSH: ${{ env.SKIP_DOCKER_PUSH }}
- name: Verify RPM signatures - name: Verify RPM signatures
@@ -294,10 +309,12 @@ jobs:
tag_and_push() { tag_and_push() {
local src="$1" img_name tag dst variant="" local src="$1" img_name tag dst variant=""
img_name="${src%%:*}" img_name="${src%%:*}"
# Client variants share a repository, so keep their tag suffixes. # Variants share a repository with their default image, so keep
# their tag suffixes. Order matters: the first matching pattern wins.
case "$src" in case "$src" in
*-rootless-ubi-amd64) variant="-rootless-ubi" ;; *-rootless-ubi-amd64) variant="-rootless-ubi" ;;
*-rootless-amd64) variant="-rootless" ;; *-rootless-amd64) variant="-rootless" ;;
*-ubi-amd64) variant="-ubi" ;;
esac esac
for tag in $(resolve_tags); do for tag in $(resolve_tags); do
dst="${img_name}:${tag}${variant}" dst="${img_name}:${tag}${variant}"
@@ -365,6 +382,24 @@ jobs:
path: dist/netbird_darwin** path: dist/netbird_darwin**
retention-days: 7 retention-days: 7
# Certify the UBI images in the Red Hat Ecosystem Catalog on stable tags.
# See redhat-certify.yml, which can also be run by hand for any released version.
redhat_certification:
name: "Red Hat"
needs: release
if: |
github.repository == 'netbirdio/netbird' &&
startsWith(github.ref, 'refs/tags/v') &&
!contains(github.ref_name, '-')
permissions:
contents: read
uses: ./.github/workflows/redhat-certify.yml
with:
component: all
version: ${{ github.ref_name }}
secrets:
PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
release_ui: release_ui:
runs-on: ubuntu-latest runs-on: ubuntu-latest
outputs: outputs:
@@ -419,12 +454,12 @@ jobs:
run: git --no-pager diff --exit-code run: git --no-pager diff --exit-code
- name: Set up Node.js - name: Set up Node.js
uses: actions/setup-node@v4 uses: actions/setup-node@v7
with: with:
node-version: '22' node-version: '22'
- name: Set up pnpm - name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0 uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with: with:
version: 11 version: 11
@@ -556,12 +591,12 @@ jobs:
run: git --no-pager diff --exit-code run: git --no-pager diff --exit-code
- name: Set up Node.js - name: Set up Node.js
uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4 uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7.0.0
with: with:
node-version: '22' node-version: '22'
- name: Set up pnpm - name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0 uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with: with:
version: 11 version: 11
@@ -653,11 +688,11 @@ jobs:
- name: check git status - name: check git status
run: git --no-pager diff --exit-code run: git --no-pager diff --exit-code
- name: Set up Node.js - name: Set up Node.js
uses: actions/setup-node@v4 uses: actions/setup-node@v7
with: with:
node-version: '22' node-version: '22'
- name: Set up pnpm - name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0 uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with: with:
version: 11 version: 11
- name: Install wails3 CLI - name: Install wails3 CLI
@@ -776,7 +811,7 @@ jobs:
run: 7z x -o"${{ github.workspace }}/NSIS_Plugins" "${{ github.workspace }}/ShellExecAsUser_amd64-Unicode.7z" run: 7z x -o"${{ github.workspace }}/NSIS_Plugins" "${{ github.workspace }}/ShellExecAsUser_amd64-Unicode.7z"
- name: Set up Go for wails3 CLI - name: Set up Go for wails3 CLI
uses: actions/setup-go@v5 uses: actions/setup-go@v6
with: with:
go-version-file: "go.mod" go-version-file: "go.mod"
cache: false cache: false
+46
View File
@@ -0,0 +1,46 @@
name: Test Homebrew cask
on:
pull_request:
paths:
- "client/ui/netbird-ui.rb.tmpl"
- ".github/scripts/test-homebrew-cask.sh"
- ".github/workflows/test-homebrew-cask.yml"
workflow_dispatch:
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }}
cancel-in-progress: true
jobs:
install-uninstall:
runs-on: macos-latest
timeout-minutes: 20
steps:
- name: Checkout code
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
with:
persist-credentials: false
- name: Clone the Homebrew tap
run: git clone https://github.com/netbirdio/homebrew-tap.git .homebrew-cask-tap
- name: Update Homebrew and install gomplate
# The runner image disables auto-update; the cask steps DSL needs Homebrew 6.0.20 or newer.
run: |
brew update
brew install gomplate
- name: Install and uninstall the cask
run: .github/scripts/test-homebrew-cask.sh
- name: Upload logs
if: always()
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a #v7.0.1
with:
name: homebrew-cask-results
path: ${{ runner.temp }}/homebrew-cask/results
if-no-files-found: ignore
+1 -1
View File
@@ -32,7 +32,7 @@ jobs:
persist-credentials: false persist-credentials: false
- name: Set up Node.js - name: Set up Node.js
uses: actions/setup-node@v4 uses: actions/setup-node@v7
with: with:
node-version: "22" node-version: "22"
+3
View File
@@ -38,4 +38,7 @@ management/server/types/testdata/
# generated by chglog in the release workflow, embedded into the RPM # generated by chglog in the release workflow, embedded into the RPM
changelog.yml changelog.yml
# generated by rpm-provides.sh, the config GoReleaser actually runs
.goreleaser.generated.yaml
.chglog.yml .chglog.yml
+98 -6
View File
@@ -60,6 +60,34 @@ builds:
- load_wgnt_from_rsrc - load_wgnt_from_rsrc
- pkcs11 - pkcs11
# Single-arch builds: nfpm provides is not templated, so the RPM splits per arch. They
# carry the PKCS#11 store like the deb build above.
- &netbird_rpm_build
id: netbird-rpm-amd64
dir: client
binary: netbird
env: [CGO_ENABLED=0]
goos: [linux]
goarch: [amd64]
ldflags:
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
mod_timestamp: "{{ .CommitTimestamp }}"
tags:
- load_wgnt_from_rsrc
- pkcs11
- <<: *netbird_rpm_build
id: netbird-rpm-arm64
goarch: [arm64]
- <<: *netbird_rpm_build
id: netbird-rpm-arm
goarch: [arm]
- <<: *netbird_rpm_build
id: netbird-rpm-386
goarch: [386]
- id: netbird-static - id: netbird-static
dir: client dir: client
binary: netbird binary: netbird
@@ -243,17 +271,22 @@ nfpms:
postinstall: "release_files/post_install.sh" postinstall: "release_files/post_install.sh"
preremove: "release_files/pre_remove.sh" preremove: "release_files/pre_remove.sh"
- maintainer: Netbird <dev@netbird.io> - &netbird_rpm
maintainer: Netbird <dev@netbird.io>
description: Netbird client. description: Netbird client.
homepage: https://netbird.io/ homepage: https://netbird.io/
license: BSD-3-Clause license: BSD-3-Clause
vendor: NetBird vendor: NetBird
id: netbird_rpm id: netbird_rpm_amd64
bindir: /usr/bin bindir: /usr/bin
builds: ids:
- netbird-pkcs11 - netbird-rpm-amd64
formats: formats:
- rpm - rpm
# Red Hat certification (RPM Version Handling) requires rpmbuild's ISA
# provide, which nfpm does not emit. The version is filled in by the release job.
provides:
- "netbird(x86-64) = @RPM_EVR@"
# The client verifies TLS to management and signal against the system trust # The client verifies TLS to management and signal against the system trust
# store. Red Hat software certification (RPM Dependency Tracking) also # store. Red Hat software certification (RPM Dependency Tracking) also
# rejects packages that declare no dependencies at all. # rejects packages that declare no dependencies at all.
@@ -283,6 +316,27 @@ nfpms:
packager: NetBird <dev@netbird.io> packager: NetBird <dev@netbird.io>
signature: signature:
key_file: '{{ if index .Env "GPG_RPM_KEY_FILE" }}{{ .Env.GPG_RPM_KEY_FILE }}{{ end }}' key_file: '{{ if index .Env "GPG_RPM_KEY_FILE" }}{{ .Env.GPG_RPM_KEY_FILE }}{{ end }}'
- <<: *netbird_rpm
id: netbird_rpm_arm64
ids:
- netbird-rpm-arm64
provides:
- "netbird(aarch-64) = @RPM_EVR@"
- <<: *netbird_rpm
id: netbird_rpm_arm
ids:
- netbird-rpm-arm
provides:
- "netbird(armv6hl-32) = @RPM_EVR@"
- <<: *netbird_rpm
id: netbird_rpm_386
ids:
- netbird-rpm-386
provides:
- "netbird(x86-32) = @RPM_EVR@"
dockers_v2: dockers_v2:
- id: netbird - id: netbird
disable: "{{ .Env.SKIP_DOCKER_PUSH }}" disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
@@ -445,7 +499,7 @@ dockers_v2:
tags: tags:
- "{{ .Version }}" - "{{ .Version }}"
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}latest{{ end }}" - "{{ if eq .Env.SKIP_PUBLISH \"false\" }}latest{{ end }}"
dockerfile: upload-server/Dockerfile dockerfile: upload-server/Dockerfile.release
platforms: platforms:
- linux/amd64 - linux/amd64
- linux/arm64 - linux/arm64
@@ -501,6 +555,41 @@ dockers_v2:
"org.opencontainers.image.revision": "{{.FullCommit}}" "org.opencontainers.image.revision": "{{.FullCommit}}"
"org.opencontainers.image.source": "{{.GitURL}}" "org.opencontainers.image.source": "{{.GitURL}}"
"maintainer": "dev@netbird.io" "maintainer": "dev@netbird.io"
- id: proxy-ubi
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
ids:
- netbird-proxy
images:
- netbirdio/reverse-proxy
- ghcr.io/netbirdio/reverse-proxy
tags:
- "{{ .Version }}-ubi"
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}ubi-latest{{ end }}"
dockerfile: proxy/Dockerfile.ubi
platforms:
- linux/amd64
- linux/arm64
build_args:
VERSION: "{{ .Version }}"
RELEASE: "{{ .Timestamp }}"
hooks:
pre:
- cmd: 'sh proxy/collect-licenses.sh "{{ .ContextDir }}/licenses" amd64 arm64'
env:
- GOOS=linux
- CGO_ENABLED=0
labels:
"org.opencontainers.image.created": "{{.Date}}"
"org.opencontainers.image.version": "{{.Version}}"
"org.opencontainers.image.revision": "{{.FullCommit}}"
"org.opencontainers.image.source": "{{.GitURL}}"
annotations:
"org.opencontainers.image.created": "{{.Date}}"
"org.opencontainers.image.title": "{{.ProjectName}}"
"org.opencontainers.image.version": "{{.Version}}"
"org.opencontainers.image.revision": "{{.FullCommit}}"
"org.opencontainers.image.source": "{{.GitURL}}"
"maintainer": "dev@netbird.io"
brews: brews:
- ids: - ids:
@@ -533,7 +622,10 @@ uploads:
- name: yum - name: yum
skip: "{{ .Env.SKIP_PUBLISH }}" skip: "{{ .Env.SKIP_PUBLISH }}"
ids: ids:
- netbird_rpm - netbird_rpm_amd64
- netbird_rpm_arm64
- netbird_rpm_arm
- netbird_rpm_386
mode: archive mode: archive
target: https://pkgs.wiretrustee.com/yum/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }} target: https://pkgs.wiretrustee.com/yum/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }}
username: dev@wiretrustee.com username: dev@wiretrustee.com
+50
View File
@@ -115,6 +115,56 @@ export NETBIRD_DOMAIN=netbird.example.com; curl -fsSL https://github.com/netbird
See a complete [architecture overview](https://docs.netbird.io/about-netbird/how-netbird-works#architecture) for details. See a complete [architecture overview](https://docs.netbird.io/about-netbird/how-netbird-works#architecture) for details.
### Reporting bugs and requesting features
NetBird uses a discussion-first workflow. Bug reports and feature requests start in
[Discussions](https://github.com/netbirdio/netbird/discussions), not as issues.
| What you want to do | Where to go |
| --- | --- |
| Report a bug, regression, or unexpected behavior | [Issue Triage](https://github.com/netbirdio/netbird/discussions/new?category=issue-triage) |
| Request a feature or share an idea | [Ideas & Feature Requests](https://github.com/netbirdio/netbird/discussions/new?category=ideas-feature-requests) |
| Ask about setup, configuration, or self-hosting | [Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support) |
| Report a security vulnerability | [Security policy](https://github.com/netbirdio/netbird/security/policy), never a public thread |
Our team and maintainers triage discussions, ask follow-up questions, check for duplicates,
and reproduce bugs. Validated reports are promoted to issues. This keeps the issue tracker a clear
answer to one question: what is the team working on.
Please search existing discussions and issues first, including closed ones. If something similar
already exists, upvote it and add your details there instead of opening a duplicate.
For bug reports, include your NetBird version, operating system, deployment type (Cloud,
self-hosted, Kubernetes, or Docker), reproduction steps, expected and actual behavior, and a debug
bundle where relevant:
```shell
netbird version
netbird status -d -A
netbird debug for 1m -A -S -U
```
`-U` uploads the bundle and prints a file key you can paste instead of attaching the archive.
`-A` anonymizes the output, which matters on a public thread. It masks most identifying details
but is not full redaction, so read the bundle before posting it. Two levels are available:
| Level | How to select | What it masks |
| --- | --- | --- |
| `default` | `-A` / `--anonymize` | Public IP addresses, IPv6 ULA addresses, MAC addresses, and domains other than `netbird.io`, `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`. IPv4 private, CGNAT, and link-local ranges are kept |
| `strict` | `--anonymize-level strict` (implies `-A`) | The above, plus IPv4 private, CGNAT, and link-local ranges, peer names in front of `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`, and WireGuard public keys. Labels under `netbird.io` are kept, since it only hosts infrastructure |
See [collecting a debug bundle](https://docs.netbird.io/help/troubleshooting-client#debug-bundle)
and the [CLI reference](https://docs.netbird.io/get-started/cli#debug-for) for details.
See [How to use Discussions, Issues, and Pull Requests](https://github.com/netbirdio/netbird/discussions/6075)
for the full workflow, or [SUPPORT.md](SUPPORT.md) for a shorter version.
### Contributing
Contributions are welcome. Read [CONTRIBUTING.md](CONTRIBUTING.md) first. NetBird works ticket
first, anything that changes behavior needs an issue the team has agreed on before you open a pull
request.
### Community projects ### Community projects
- [NetBird installer script](https://github.com/physk/netbird-installer) - [NetBird installer script](https://github.com/physk/netbird-installer)
- [netbird-tui](https://github.com/n0pashkov/netbird-tui) - terminal UI for managing NetBird peers, routes, and settings - [netbird-tui](https://github.com/n0pashkov/netbird-tui) - terminal UI for managing NetBird peers, routes, and settings
+121
View File
@@ -0,0 +1,121 @@
# Getting help with NetBird
Where to go depends on what you need. If you are not sure, start with
[Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support)
and we will move it.
## Before you post
1. Search existing [discussions](https://github.com/netbirdio/netbird/discussions) and
[issues](https://github.com/netbirdio/netbird/issues), including closed ones.
2. Check the [documentation](https://docs.netbird.io) and the troubleshooting guides for
[clients](https://docs.netbird.io/help/troubleshooting-client) and
[self-hosted deployments](https://docs.netbird.io/selfhosted/troubleshooting).
3. Remove or anonymize sensitive information from logs, screenshots, and configuration.
If a discussion already covers your problem, upvote it and add your details there rather than
opening a duplicate. Extra reproduction detail, affected versions, and deployment notes are
useful even on an existing thread.
## Community support
Free, for everyone. Covers the NetBird client, open source self-hosted deployments, and general
questions.
| What you want to do | Where to go |
| --- | --- |
| Report a bug, regression, or unexpected behavior | [Issue Triage](https://github.com/netbirdio/netbird/discussions/new?category=issue-triage) |
| Request a feature or share an idea | [Ideas & Feature Requests](https://github.com/netbirdio/netbird/discussions/new?category=ideas-feature-requests) |
| Ask about setup, configuration, or self-hosting | [Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support) |
| Chat with the community | [Slack](https://docs.netbird.io/slack-url) |
## Paid support
For NetBird Cloud customers and commercial-license self-hosted deployments, covering the
dashboard, control plane, billing, and subscriptions, see
[reporting bugs and issues](https://docs.netbird.io/help/report-bug-issues).
## Security
Do not report security vulnerabilities in public issues or discussions, and do not post secrets,
private keys, internal hostnames, or sensitive logs. Use the
[security policy](https://github.com/netbirdio/netbird/security/policy).
## What makes a report we can act on
For a bug, the most useful reports include:
- NetBird version, and component versions where applicable
- Operating system or environment
- Deployment type: NetBird Cloud, self-hosted, Kubernetes, Docker, or local development
- Current behavior and expected behavior
- The smallest set of steps that reproduces the problem
- Logs, status output, screenshots, or a debug bundle when relevant
- Whether this worked before, and the last known working version
For client reports, these commands usually give us what we need:
```shell
netbird version
netbird status -d -A
netbird debug for 1m -A -S -U
```
`-A` (`--anonymize`) replaces sensitive values consistently across every file in the bundle, so
it stays readable while masking most identifying details. It is not a guarantee of full redaction:
internal address ranges survive at the default level, and interface names, indexes, MTUs, and
flags are never anonymized. Read the bundle before posting it publicly. Two levels are
available:
| Level | How to select | What it masks |
| --- | --- | --- |
| `default` | `-A` / `--anonymize`, or `--anonymize-level default` | Public IP addresses, IPv6 ULA addresses, MAC addresses, and domains other than `netbird.io`, `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`. IPv4 private, CGNAT, and link-local ranges are kept, and interface names are not anonymized |
| `strict` | `--anonymize-level strict` (implies `-A`) | The above, plus IPv4 private, CGNAT, and link-local ranges, peer names in front of `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`, and WireGuard public keys. Labels under `netbird.io` are kept, since it only hosts infrastructure |
Use `strict` when internal addressing or peer naming is itself sensitive. Either way, private
keys and SSH keys are never included, and the packet capture (`capture.pcap`) is left out of
anonymized bundles because it holds raw decrypted packets.
`-U` (`--upload-bundle`) uploads the bundle and returns a file key you can paste into the thread
instead of attaching an archive. Retention is controlled by the upload service; check its policy
before uploading, and configure cleanup for self-hosted deployments.
For more detail, see [troubleshooting client issues](https://docs.netbird.io/help/troubleshooting-client),
which explains [what a debug bundle contains](https://docs.netbird.io/help/troubleshooting-client#debug-bundle),
and the [CLI reference](https://docs.netbird.io/get-started/cli#debug-for).
Intermittent problems are still worth reporting. They just need enough detail to investigate:
trigger, frequency, timing, timestamps, and any related logs.
For a feature request, describe the problem before the solution: what you are trying to
accomplish, who is affected and how often, why the current behavior or workaround is not enough,
and what you would like to see instead.
## What happens after you post
Our team, maintainers, or community members may ask for missing details, link related
threads, merge duplicates, move your post to a better category, or try to reproduce the problem.
Not every discussion becomes an issue. Some are answered in Q&A, some turn out to be
configuration problems, and some need more information before engineering can act. A
well-answered discussion is still a useful outcome.
When a report is confirmed and actionable, a maintainer opens a validated issue linked back to
the discussion, in whichever repository the fix belongs to. You do not need to know which
repository that is. Routing is part of triage.
## A note on issues
Issues in this repository are maintainer-curated work items. Every open issue is something a
maintainer or contributor can pick up and act on. Issues opened without a linked validated
discussion may be closed and redirected here.
Maintainers can still open issues directly for work found internally, such as regressions caught
during development, planned maintenance, or release blockers.
## Related reading
- [How to use Discussions, Issues, and Pull Requests](https://github.com/netbirdio/netbird/discussions/6075)
- [Moving to a discussion-first approach](https://github.com/netbirdio/netbird/discussions/6074)
- [CONTRIBUTING.md](CONTRIBUTING.md) for opening pull requests
- [CODE_OF_CONDUCT.md](CODE_OF_CONDUCT.md)
+6 -6
View File
@@ -110,12 +110,12 @@ Two roles delegate Agent Network access without account-admin rights:
read-only users, groups, peers, and account info (needed to build policies). read-only users, groups, peers, and account info (needed to build policies).
Nothing else in the account. Nothing else in the account.
- **`usage_viewer`** — the regular User baseline plus read on - **`usage_viewer`** — the regular User baseline plus read on
`agent_network.usage` (the aggregated usage and cost overview) and read-only `agent_network.usage` (the aggregated usage and cost overview) and
access to the resources the usage filters resolve against: users, groups, `agent_network.logs` (the account-wide request-level access logs, which can
peers, and the provider list (connection config redacted — no upstream URLs contain captured prompts), and read-only access to the resources those
or operator-supplied header values). No policies, and no account-wide filters resolve against: users, groups, peers, and the provider list
request-level access logs; like any caller, it still reads its own requests (connection config redacted — no upstream URLs or operator-supplied header
through the self-scoped endpoints below. values). No policies, guardrails, budgets, or settings.
Every authenticated user, regardless of role, can read the caller-scoped Every authenticated user, regardless of role, can read the caller-scoped
self-service endpoint `GET /api/agent-network/agent-config` (the endpoint, providers, self-service endpoint `GET /api/agent-network/agent-config` (the endpoint, providers,
+52 -33
View File
@@ -3,56 +3,75 @@ package base62
import ( import (
"fmt" "fmt"
"math" "math"
"strings"
) )
const ( const (
alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
base = uint32(len(alphabet)) base = uint32(len(alphabet))
maxBase62Digits = 6 // max number of digits required to encode MaxUint32
) )
var (
ErrEmptyString = fmt.Errorf("empty string")
ErrInvalidChar = fmt.Errorf("invalid character")
ErrOverflow = fmt.Errorf("integer overflow")
)
// Fixed-size arrays have better performance and lower memory overhead compared to maps for small static sets of data
var charToIndex [123]int8 // Assuming ASCII from '\0' - 'z'
func init() {
for i := range charToIndex {
charToIndex[i] = -1
}
for i, c := range alphabet {
charToIndex[c] = int8(i)
}
}
// Encode encodes a uint32 value to a base62 string. // Encode encodes a uint32 value to a base62 string.
func Encode(num uint32) string { // The returned string will be between 1-6 characters long.
if num == 0 { func Encode(n uint32) string {
return string(alphabet[0]) if n < base {
return string(alphabet[n])
}
// avoid dynamic memory usage for small, fixed size data
buf := [maxBase62Digits]byte{}
idx := len(buf)
for n > 0 {
idx--
buf[idx] = alphabet[n%base]
n /= base
} }
var encoded strings.Builder return string(buf[idx:])
for num > 0 {
remainder := num % base
encoded.WriteByte(alphabet[remainder])
num /= base
}
// Reverse the encoded string
encodedString := encoded.String()
reversed := reverse(encodedString)
return reversed
} }
// Decode decodes a base62 string to a uint32 value. // Decode decodes a base62 string to a uint32 value.
// Returns an error if the input string is empty, contains invalid characters,
// or would result in integer overflow.
func Decode(encoded string) (uint32, error) { func Decode(encoded string) (uint32, error) {
if len(encoded) == 0 {
return 0, ErrEmptyString
}
var decoded uint32 var decoded uint32
strLen := len(encoded) for _, char := range encoded {
index := int8(-1)
for i, char := range encoded { if int(char) < len(charToIndex) {
index := strings.IndexRune(alphabet, char) index = charToIndex[char]
}
if index < 0 { if index < 0 {
return 0, fmt.Errorf("invalid character: %c", char) return 0, fmt.Errorf("%w: %c", ErrInvalidChar, char)
}
// Add overflow check when calculating the decoded value to prevent silent overflow of uint32
if decoded > (math.MaxUint32-uint32(index))/base {
return 0, fmt.Errorf("%w: %s", ErrOverflow, encoded)
} }
decoded += uint32(index) * uint32(math.Pow(float64(base), float64(strLen-i-1))) decoded = decoded*base + uint32(index)
} }
return decoded, nil return decoded, nil
} }
// Reverse a string.
func reverse(s string) string {
runes := []rune(s)
for i, j := 0, len(runes)-1; i < j; i, j = i+1, j-1 {
runes[i], runes[j] = runes[j], runes[i]
}
return string(runes)
}
+50 -14
View File
@@ -1,31 +1,67 @@
package base62 package base62
import ( import (
"errors"
"math"
"testing" "testing"
) )
func TestEncodeDecode(t *testing.T) { func TestEncodeDecode(t *testing.T) {
tests := []struct { testCases := []struct {
num uint32 input uint32
expected string
}{ }{
{0}, {0, "0"},
{1}, {1, "1"},
{42}, {5, "5"},
{12345}, {9, "9"},
{99999}, {10, "A"},
{123456789}, {42, "g"},
{61, "z"},
{62, "10"},
{'0', "m"},
{'9', "v"},
{'A', "13"},
{'Z', "1S"},
{'a', "1Z"},
{'z', "1y"},
{99999, "Q0t"},
{12345, "3D7"},
{123456789, "8M0kX"},
{math.MaxUint32, "4gfFC3"},
} }
for _, tt := range tests { for _, tc := range testCases {
encoded := Encode(tt.num) encoded := Encode(tc.input)
if encoded != tc.expected {
t.Errorf("Encode(%d) = %s; want %s", tc.input, encoded, tc.expected)
}
decoded, err := Decode(encoded) decoded, err := Decode(encoded)
if err != nil { if err != nil {
t.Errorf("Decode error: %v", err) t.Errorf("Expected error nil, got %v", err)
} }
if decoded != tt.num { if decoded != tc.input {
t.Errorf("Decode(%v) = %v, want %v", encoded, decoded, tt.num) t.Errorf("Decode(%v) = %v, want %v", encoded, decoded, tc.input)
} }
} }
} }
// Decode handles empty string input with appropriate error
func TestDecodeEmptyString(t *testing.T) {
if _, err := Decode(""); !errors.Is(err, ErrEmptyString) {
t.Errorf("Expected error %v, got %v", ErrEmptyString, err)
}
}
func TestDecodeOverflow(t *testing.T) {
if _, err := Decode("4gfFC4"); !errors.Is(err, ErrOverflow) {
t.Errorf("Expected error %v, got %v", ErrOverflow, err)
}
}
func TestDecodeInvalid(t *testing.T) {
if _, err := Decode("/"); !errors.Is(err, ErrInvalidChar) {
t.Errorf("Expected error %v, got %v", ErrInvalidChar, err)
}
}
+31 -11
View File
@@ -104,8 +104,7 @@ type Client struct {
stateChangeMu sync.Mutex stateChangeMu sync.Mutex
stateChangeSubID string stateChangeSubID string
eventSub *peer.EventSubscription // Closed to stop the watch goroutine from delivering buffered ticks to a
// Closed to stop the watch goroutines from delivering buffered items to a
// listener that has been removed or replaced. See stopStateChangeWatchLocked. // listener that has been removed or replaced. See stopStateChangeWatchLocked.
stateChangeDone chan struct{} stateChangeDone chan struct{}
@@ -213,6 +212,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
internal.WithNetEvents(c.netMgr)) internal.WithNetEvents(c.netMgr))
c.setState(cfg, cacheDir, cfgFile, connectClient) c.setState(cfg, cacheDir, cfgFile, connectClient)
connectClient.SetSyncResponsePersistence(true)
// This path runs the interactive SSO flow, so reaching here means the peer // This path runs the interactive SSO flow, so reaching here means the peer
// is authenticated again — release the latch Status() reports from. Clear // is authenticated again — release the latch Status() reports from. Clear
// only once the fresh connect client is installed: until then Status() // only once the fresh connect client is installed: until then Status()
@@ -256,6 +256,7 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
internal.WithNetEvents(c.netMgr)) internal.WithNetEvents(c.netMgr))
c.setState(cfg, cacheDir, cfgFile, connectClient) c.setState(cfg, cacheDir, cfgFile, connectClient)
connectClient.SetSyncResponsePersistence(true)
return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir) return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir)
} }
@@ -327,6 +328,19 @@ func (c *Client) NotifyNetworkChange() {
// or "strict"; strict also anonymizes internal IP ranges, peer names, and // or "strict"; strict also anonymizes internal IP ranges, peer names, and
// WireGuard public keys, and implies anonymize. // WireGuard public keys, and implies anonymize.
func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) { func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) {
return c.debugBundle(platformFiles, anonymize, anonymizeLevel, true)
}
// DebugBundleFile generates a debug bundle and returns the path of the zip in
// the cache directory instead of uploading it, so the app can hand the file to
// the user for inspection. The caller owns the file and removes it once done;
// the stale-bundle cleanup of later runs removes it only after a day.
// anonymize and anonymizeLevel behave as in DebugBundle.
func (c *Client) DebugBundleFile(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) {
return c.debugBundle(platformFiles, anonymize, anonymizeLevel, false)
}
func (c *Client) debugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string, upload bool) (string, error) {
cfg, cacheDir, cc := c.stateSnapshot() cfg, cacheDir, cc := c.stateSnapshot()
// If the engine hasn't been started, load config from disk // If the engine hasn't been started, load config from disk
@@ -342,6 +356,11 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
cacheDir = platformFiles.CacheDir() cacheDir = platformFiles.CacheDir()
} }
// Clear what an interrupted earlier run may have left in the cache before
// adding to it. Remote debug jobs write to the same directory, so anything
// younger than an hour is treated as possibly still in use.
debug.RemoveStaleBundles(cacheDir, time.Hour)
deps := debug.GeneratorDependencies{ deps := debug.GeneratorDependencies{
InternalConfig: cfg, InternalConfig: cfg,
StatusRecorder: c.recorder, StatusRecorder: c.recorder,
@@ -379,6 +398,9 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
if err != nil { if err != nil {
return "", fmt.Errorf("generate debug bundle: %w", err) return "", fmt.Errorf("generate debug bundle: %w", err)
} }
if !upload {
return debug.ExportBundle(path)
}
defer func() { defer func() {
if err := os.Remove(path); err != nil { if err := os.Remove(path); err != nil {
log.Errorf("failed to remove debug bundle file: %v", err) log.Errorf("failed to remove debug bundle file: %v", err)
@@ -475,6 +497,7 @@ func (c *Client) Networks() *NetworkArray {
routesMap := routeManager.GetClientRoutesWithNetID() routesMap := routeManager.GetClientRoutesWithNetID()
v6Merged := route.V6ExitMergeSet(routesMap) v6Merged := route.V6ExitMergeSet(routesMap)
resolvedDomains := c.recorder.GetResolvedDomainsStates() resolvedDomains := c.recorder.GetResolvedDomainsStates()
activeRoutePeers := c.recorder.GetActiveRoutePeers()
networkArray := &NetworkArray{ networkArray := &NetworkArray{
items: make([]Network, 0), items: make([]Network, 0),
@@ -488,7 +511,7 @@ func (c *Client) Networks() *NetworkArray {
continue continue
} }
network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged) network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged, activeRoutePeers)
if network == nil { if network == nil {
continue continue
} }
@@ -497,14 +520,14 @@ func (c *Client) Networks() *NetworkArray {
return networkArray return networkArray
} }
func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}) *Network { func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}, activeRoutePeers map[route.HAUniqueID]string) *Network {
r := routes[0] r := routes[0]
netStr := r.Network.String() netStr := r.Network.String()
if r.IsDynamic() { if r.IsDynamic() {
netStr = r.Domains.SafeString() netStr = r.Domains.SafeString()
} }
routePeer, err := c.findBestRoutePeer(routes) routePeer, err := c.findBestRoutePeer(routes, activeRoutePeers)
if err != nil { if err != nil {
log.Errorf("could not get peer info for route %s: %v", id, err) log.Errorf("could not get peer info for route %s: %v", id, err)
return nil return nil
@@ -528,12 +551,9 @@ func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bo
// findBestRoutePeer returns the peer actively routing traffic for the given // findBestRoutePeer returns the peer actively routing traffic for the given
// HA route group. Falls back to the first connected peer, then the first peer. // HA route group. Falls back to the first connected peer, then the first peer.
func (c *Client) findBestRoutePeer(routes []*route.Route) (peer.State, error) { func (c *Client) findBestRoutePeer(routes []*route.Route, activeRoutePeers map[route.HAUniqueID]string) (peer.State, error) {
netStr := routes[0].Network.String() if peerKey, ok := activeRoutePeers[routes[0].GetHAUniqueID()]; ok {
if p, err := c.recorder.GetPeer(peerKey); err == nil {
fullStatus := c.recorder.GetFullStatus()
for _, p := range fullStatus.Peers {
if _, ok := p.GetRoutes()[netStr]; ok {
return p, nil return p, nil
} }
} }
+10 -81
View File
@@ -6,13 +6,8 @@ import (
"context" "context"
"fmt" "fmt"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal" "github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/internal/auth" "github.com/netbirdio/netbird/client/internal/auth"
"github.com/netbirdio/netbird/client/internal/auth/sessionwatch"
"github.com/netbirdio/netbird/client/internal/peer"
cProto "github.com/netbirdio/netbird/client/proto"
) )
// StateChangeListener receives client state notifications. // StateChangeListener receives client state notifications.
@@ -21,16 +16,11 @@ import (
// changed: connection state, the run-loop status label (e.g. NeedsLogin) or // changed: connection state, the run-loop status label (e.g. NeedsLogin) or
// the session deadline. It mirrors the daemon's SubscribeStatus stream // the session deadline. It mirrors the daemon's SubscribeStatus stream
// trigger — on each signal the consumer pulls the fresh values via // trigger — on each signal the consumer pulls the fresh values via
// Status() / SessionExpiresAtUnix(). // Status() / SessionExpiresAtUnix(). The engine arms no expiry-warning
// // timers on Android; the app schedules the warnings from the deadline it
// OnSessionExpiring forwards the engine's session-expiry warnings, fired at // reads here.
// sessionwatch.WarningLead before the deadline and again at FinalWarningLead
// (finalWarning true). The second one is suppressed when the user dismissed
// the first via DismissSessionWarning. The daemon turns the same events into
// its tray notification.
type StateChangeListener interface { type StateChangeListener interface {
OnStateChanged() OnStateChanged()
OnSessionExpiring(expiresAtUnix int64, leadMinutes int64, finalWarning bool)
} }
// Status returns the connect run-loop's status label — the same value the // Status returns the connect run-loop's status label — the same value the
@@ -110,11 +100,11 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) {
return return
} }
// Both subscriptions are buffered (one pending tick, ten pending events), // The subscription is buffered (one pending tick), so unsubscribing is
// so unsubscribing is not enough to stop callbacks: the loops would drain // not enough to stop callbacks: the loop would drain what is already
// what is already queued and deliver it to a listener the caller has // queued and deliver it to a listener the caller has already removed or
// already removed or replaced. Gate every callback on this registration's // replaced. Gate every callback on this registration's own signal, which
// own signal, which is closed before unsubscribing. // is closed before unsubscribing.
done := make(chan struct{}) done := make(chan struct{})
c.stateChangeDone = done c.stateChangeDone = done
@@ -133,9 +123,6 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) {
listener.OnStateChanged() listener.OnStateChanged()
} }
}() }()
c.eventSub = c.recorder.SubscribeToEvents()
go watchSessionWarnings(c.eventSub, listener, done)
} }
// RemoveStateChangeListener unregisters the state notification listener. // RemoveStateChangeListener unregisters the state notification listener.
@@ -145,21 +132,6 @@ func (c *Client) RemoveStateChangeListener() {
c.stopStateChangeWatchLocked() c.stopStateChangeWatchLocked()
} }
// DismissSessionWarning records the user's "Dismiss" on the first expiry
// warning and suppresses the final one for the current deadline. A refreshed
// deadline re-arms both. No-op while the engine is not running.
func (c *Client) DismissSessionWarning() {
cc := c.getConnectClient()
if cc == nil {
return
}
engine := cc.Engine()
if engine == nil {
return
}
engine.DismissSessionWarning()
}
// ExtendAuthSession runs the interactive SSO flow to obtain a fresh JWT and // ExtendAuthSession runs the interactive SSO flow to obtain a fresh JWT and
// asks the management server to extend the session deadline. The tunnel is // asks the management server to extend the session deadline. The tunnel is
// untouched: no resync, no reconnect. Async; the result arrives on the // untouched: no resync, no reconnect. Async; the result arrives on the
@@ -201,8 +173,8 @@ func (c *Client) CancelExtendAuthSession() {
} }
func (c *Client) stopStateChangeWatchLocked() { func (c *Client) stopStateChangeWatchLocked() {
// Signal first, unsubscribe second: closing the channels only stops new // Signal first, unsubscribe second: closing the channel only stops new
// items, and the loops would still hand whatever is buffered to a listener // items, and the loop would still hand whatever is buffered to a listener
// that is no longer registered. // that is no longer registered.
if c.stateChangeDone != nil { if c.stateChangeDone != nil {
close(c.stateChangeDone) close(c.stateChangeDone)
@@ -212,49 +184,6 @@ func (c *Client) stopStateChangeWatchLocked() {
c.recorder.UnsubscribeFromStateChanges(c.stateChangeSubID) c.recorder.UnsubscribeFromStateChanges(c.stateChangeSubID)
c.stateChangeSubID = "" c.stateChangeSubID = ""
} }
if c.eventSub != nil {
// Closes the channel, which ends watchSessionWarnings.
c.recorder.UnsubscribeFromEvents(c.eventSub)
c.eventSub = nil
}
}
// watchSessionWarnings forwards the engine's session-expiry warnings to the
// listener. The event stream also carries unrelated traffic — network-map
// updates on every sync, DNS and route errors — so everything but an
// AUTHENTICATION event carrying the session-warning marker is dropped. Exits
// when the subscription is closed by UnsubscribeFromEvents, or earlier when
// done is closed — the stream buffers up to ten events, and a deregistered
// listener must not receive the ones already queued.
func watchSessionWarnings(sub *peer.EventSubscription, listener StateChangeListener, done <-chan struct{}) {
for ev := range sub.Events() {
select {
case <-done:
return
default:
}
if ev.GetCategory() != cProto.SystemEvent_AUTHENTICATION {
continue
}
meta := ev.GetMetadata()
if meta[sessionwatch.MetaSessionWarning] != "true" {
// Other AUTHENTICATION events exist (e.g. a deadline rejected as
// out of range); they carry no warning marker.
continue
}
deadline, err := sessionwatch.ParseExpiresAt(meta[sessionwatch.MetaSessionExpiresAt])
if err != nil {
log.Warnf("session warning event with unparsable deadline: %v", err)
continue
}
lead, err := sessionwatch.ParseLeadMinutes(meta[sessionwatch.MetaSessionLeadMinutes])
if err != nil {
// Informational only — the deadline above is what drives the UI.
lead = 0
}
listener.OnSessionExpiring(deadline.Unix(), int64(lead),
meta[sessionwatch.MetaSessionFinal] == "true")
}
} }
func (c *Client) beginExtend() (context.Context, error) { func (c *Client) beginExtend() (context.Context, error) {
+2
View File
@@ -31,6 +31,8 @@ const (
// PasswordRequiredMarker tells Java to prompt for a password and retry. It is // PasswordRequiredMarker tells Java to prompt for a password and retry. It is
// a string because gomobile flattens errors to their message, so a sentinel // a string because gomobile flattens errors to their message, so a sentinel
// value would not survive the binding. // value would not survive the binding.
//
//nolint:gosec // G101 false positive: a sentinel marker, not a credential
const PasswordRequiredMarker = "netbird-ssh-password-required" const PasswordRequiredMarker = "netbird-ssh-password-required"
// HostKeyUnknownMarker tells Java to show the fingerprint and, on confirmation, // HostKeyUnknownMarker tells Java to show the fingerprint and, on confirmation,
+58 -26
View File
@@ -23,7 +23,10 @@ import (
"github.com/netbirdio/netbird/version" "github.com/netbirdio/netbird/version"
) )
const errCloseConnection = "Failed to close connection: %v" const (
errCloseConnection = "Failed to close connection: %v"
noUpDownFlag = "no-updown"
)
var ( var (
logFileCount uint32 logFileCount uint32
@@ -257,13 +260,14 @@ func runForDuration(cmd *cobra.Command, args []string) error {
} }
stateWasDown := stat.Status != string(internal.StatusConnected) && stat.Status != string(internal.StatusConnecting) stateWasDown := stat.Status != string(internal.StatusConnected) && stat.Status != string(internal.StatusConnecting)
noUpDown, _ := cmd.Flags().GetBool(noUpDownFlag)
initialLogLevel, err := client.GetLogLevel(cmd.Context(), &proto.GetLogLevelRequest{}) initialLogLevel, err := client.GetLogLevel(cmd.Context(), &proto.GetLogLevelRequest{})
if err != nil { if err != nil {
return fmt.Errorf("failed to get log level: %v", status.Convert(err).Message()) return fmt.Errorf("failed to get log level: %v", status.Convert(err).Message())
} }
if stateWasDown { if stateWasDown && !noUpDown {
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil { if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message()) cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
} else { } else {
@@ -284,34 +288,20 @@ func runForDuration(cmd *cobra.Command, args []string) error {
} }
needsRestoreUp := false needsRestoreUp := false
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil { if noUpDown {
cmd.PrintErrf("Failed to bring service down: %v\n", status.Convert(err).Message()) enableSyncResponsePersistence(cmd, client)
} else { } else {
needsRestoreUp = !stateWasDown needsRestoreUp = restartDaemon(cmd, client, stateWasDown)
cmd.Println("netbird down")
} }
time.Sleep(1 * time.Second)
// Enable sync response persistence before bringing the service up
if _, err := client.SetSyncResponsePersistence(cmd.Context(), &proto.SetSyncResponsePersistenceRequest{
Enabled: true,
}); err != nil {
cmd.PrintErrf("Failed to enable sync response persistence: %v\n", status.Convert(err).Message())
}
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
} else {
needsRestoreUp = false
cmd.Println("netbird up")
}
time.Sleep(3 * time.Second)
cpuProfilingStarted := false cpuProfilingStarted := false
if _, err := client.StartCPUProfile(cmd.Context(), &proto.StartCPUProfileRequest{}); err != nil { if _, err := client.StartCPUProfile(cmd.Context(), &proto.StartCPUProfileRequest{}); err != nil {
cmd.PrintErrf("Failed to start CPU profiling: %v\n", err) if msg := status.Convert(err).Message(); strings.Contains(msg, "already in progress") {
cmd.PrintErrln("CPU profiling is already running (started with `netbird debug cpu start`). " +
"It is left running and is included in a bundle created after `netbird debug cpu stop`.")
} else {
cmd.PrintErrf("Failed to start CPU profiling: %v\n", msg)
}
} else { } else {
cpuProfilingStarted = true cpuProfilingStarted = true
defer func() { defer func() {
@@ -401,7 +391,7 @@ func runForDuration(cmd *cobra.Command, args []string) error {
} }
} }
if stateWasDown { if stateWasDown && !noUpDown {
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil { if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
cmd.PrintErrf("Failed to restore service down state: %v\n", status.Convert(err).Message()) cmd.PrintErrf("Failed to restore service down state: %v\n", status.Convert(err).Message())
} else { } else {
@@ -458,6 +448,47 @@ func setSyncResponsePersistence(cmd *cobra.Command, args []string) error {
return nil return nil
} }
// enableSyncResponsePersistence asks the daemon to keep the latest sync
// response so the bundle carries the network map. With a running daemon only
// syncs received after the call are kept.
func enableSyncResponsePersistence(cmd *cobra.Command, client proto.DaemonServiceClient) {
if _, err := client.SetSyncResponsePersistence(cmd.Context(), &proto.SetSyncResponsePersistenceRequest{
Enabled: true,
}); err != nil {
cmd.PrintErrf("Failed to enable sync response persistence: %v\n", status.Convert(err).Message())
}
}
// restartDaemon cycles the daemon down and up with sync response persistence
// enabled so the bundle carries the network map. It reports whether the
// daemon was left down although it was running before, so the caller can
// bring it back up.
func restartDaemon(cmd *cobra.Command, client proto.DaemonServiceClient, stateWasDown bool) bool {
needsRestoreUp := false
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
cmd.PrintErrf("Failed to bring service down: %v\n", status.Convert(err).Message())
} else {
needsRestoreUp = !stateWasDown
cmd.Println("netbird down")
}
time.Sleep(1 * time.Second)
// Enable sync response persistence before bringing the service up
enableSyncResponsePersistence(cmd, client)
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
} else {
needsRestoreUp = false
cmd.Println("netbird up")
}
time.Sleep(3 * time.Second)
return needsRestoreUp
}
func waitForDurationOrCancel(ctx context.Context, duration time.Duration, cmd *cobra.Command) error { func waitForDurationOrCancel(ctx context.Context, duration time.Duration, cmd *cobra.Command) error {
ticker := time.NewTicker(1 * time.Second) ticker := time.NewTicker(1 * time.Second)
defer ticker.Stop() defer ticker.Stop()
@@ -546,4 +577,5 @@ func init() {
forCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle") forCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle")
forCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root") forCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root")
forCmd.Flags().Bool("capture", false, "Capture packets during the debug duration and include in bundle") forCmd.Flags().Bool("capture", false, "Capture packets during the debug duration and include in bundle")
forCmd.Flags().Bool(noUpDownFlag, false, "Keep the daemon running instead of bringing it down and up before collecting. The bundle only includes the network map if a sync arrives during the run")
} }
+83
View File
@@ -0,0 +1,83 @@
package cmd
import (
"fmt"
log "github.com/sirupsen/logrus"
"github.com/spf13/cobra"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/proto"
)
var debugCPUCmd = &cobra.Command{
Use: "cpu",
Short: "Profile the daemon's CPU usage",
Long: `Starts and stops CPU profiling in the running daemon without restarting it.
The profile is included in the next debug bundle as cpu.prof.
Profiling is not time limited: it keeps running, and keeps costing CPU, until
"netbird debug cpu stop" is run.`,
}
var debugCPUStartCmd = &cobra.Command{
Use: "start",
Short: "Start CPU profiling in the daemon",
Example: " netbird debug cpu start",
Args: cobra.NoArgs,
RunE: debugCPUStart,
}
var debugCPUStopCmd = &cobra.Command{
Use: "stop",
Short: "Stop CPU profiling in the daemon",
Long: `Stops CPU profiling. The captured profile stays in the daemon until the next
debug bundle is created, which includes it as cpu.prof.`,
Example: " netbird debug cpu stop && netbird debug bundle",
Args: cobra.NoArgs,
RunE: debugCPUStop,
}
func debugCPUStart(cmd *cobra.Command, _ []string) error {
conn, err := getClient(cmd)
if err != nil {
return err
}
defer func() {
if err := conn.Close(); err != nil {
log.Errorf(errCloseConnection, err)
}
}()
if _, err := proto.NewDaemonServiceClient(conn).StartCPUProfile(cmd.Context(), &proto.StartCPUProfileRequest{}); err != nil {
return fmt.Errorf("start CPU profiling: %v", status.Convert(err).Message())
}
cmd.Println("CPU profiling started and runs until stopped. Run `netbird debug cpu stop` and then `netbird debug bundle` to collect it.")
return nil
}
func debugCPUStop(cmd *cobra.Command, _ []string) error {
conn, err := getClient(cmd)
if err != nil {
return err
}
defer func() {
if err := conn.Close(); err != nil {
log.Errorf(errCloseConnection, err)
}
}()
if _, err := proto.NewDaemonServiceClient(conn).StopCPUProfile(cmd.Context(), &proto.StopCPUProfileRequest{}); err != nil {
return fmt.Errorf("stop CPU profiling: %v", status.Convert(err).Message())
}
cmd.Println("CPU profiling stopped. Run `netbird debug bundle` to include cpu.prof.")
return nil
}
func init() {
debugCPUCmd.AddCommand(debugCPUStartCmd)
debugCPUCmd.AddCommand(debugCPUStopCmd)
debugCmd.AddCommand(debugCPUCmd)
}
+164
View File
@@ -0,0 +1,164 @@
package cmd
import (
"bytes"
"context"
"os/user"
"strings"
"testing"
"github.com/spf13/cobra"
"github.com/spf13/pflag"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/internal/profilemanager"
)
// startDebugTestDaemon starts an in-process daemon with an isolated profile
// directory and returns the address the CLI should dial.
func startDebugTestDaemon(t *testing.T) string {
t.Helper()
tempDir := t.TempDir()
origDefaultProfileDir := profilemanager.DefaultConfigPathDir
origActiveProfileStatePath := profilemanager.ActiveProfileStatePath
origConfigDirOverride := profilemanager.ConfigDirOverride
origDaemonAddr := daemonAddr
t.Cleanup(func() {
profilemanager.DefaultConfigPathDir = origDefaultProfileDir
profilemanager.ActiveProfileStatePath = origActiveProfileStatePath
profilemanager.ConfigDirOverride = origConfigDirOverride
daemonAddr = origDaemonAddr
})
profilemanager.DefaultConfigPathDir = tempDir
profilemanager.ActiveProfileStatePath = tempDir + "/active_profile.json"
profilemanager.ConfigDirOverride = tempDir
currUser, err := user.Current()
require.NoError(t, err)
sm := profilemanager.ServiceManager{}
created, err := sm.AddProfile("test1", currUser.Username)
require.NoError(t, err)
require.NoError(t, sm.SetActiveProfileState(&profilemanager.ActiveProfileState{
ID: created.ID,
Username: currUser.Username,
}))
ctx, cancel := context.WithCancel(internal.CtxInitState(context.Background()))
srv, lis := startClientDaemon(t, ctx, "", tempDir+"/config.json")
t.Cleanup(func() {
cancel()
srv.Stop()
})
return "tcp://" + lis.Addr().String()
}
// runDebugCmd runs `netbird debug <args>` against the daemon at addr and
// returns everything the command printed.
func runDebugCmd(addr string, args ...string) (string, error) {
daemonAddr = addr
var out bytes.Buffer
rootCmd.SetOut(&out)
rootCmd.SetErr(&out)
rootCmd.SetArgs(append(append([]string{"debug"}, args...), "--daemon-addr", addr, "--log-file", ""))
err := rootCmd.Execute()
rootCmd.SetOut(nil)
rootCmd.SetErr(nil)
rootCmd.SetArgs(nil)
resetFlags(rootCmd)
return out.String(), err
}
// resetFlags puts every flag of the command and its subcommands back to its
// default so a value parsed in one run does not leak into the next in-process
// execution.
func resetFlags(cmd *cobra.Command) {
reset := func(f *pflag.Flag) {
// Set appends to a slice flag and would parse the "[a,b]" default
// text as elements, so slices are replaced instead.
if sv, ok := f.Value.(pflag.SliceValue); ok {
var def []string
if trimmed := strings.Trim(f.DefValue, "[]"); trimmed != "" {
def = strings.Split(trimmed, ",")
}
_ = sv.Replace(def)
} else {
_ = f.Value.Set(f.DefValue)
}
f.Changed = false
}
cmd.Flags().VisitAll(reset)
cmd.PersistentFlags().VisitAll(reset)
// Commands pin their writers to the buffer of the run that first used
// them, so a later run would print into the old buffer.
cmd.SetOut(nil)
cmd.SetErr(nil)
for _, sub := range cmd.Commands() {
resetFlags(sub)
}
}
// TestResetFlagsSliceDefault guards against Set("[]") on slice flags, which
// stores a literal "[]" element instead of the empty default.
func TestResetFlagsSliceDefault(t *testing.T) {
cmd := &cobra.Command{Use: "x"}
var env, withDefault []string
cmd.Flags().StringSliceVar(&env, "env", nil, "")
cmd.Flags().StringSliceVar(&withDefault, "names", []string{"a", "b"}, "")
require.NoError(t, cmd.Flags().Parse([]string{"--env", "K=V", "--names", "c"}))
resetFlags(cmd)
assert.Empty(t, env, "slice flag with no default must reset to empty")
assert.Equal(t, []string{"a", "b"}, withDefault, "slice flag must reset to its default")
}
func TestDebugCPUStartStop(t *testing.T) {
addr := startDebugTestDaemon(t)
run := func(args ...string) error {
_, err := runDebugCmd(addr, append([]string{"cpu"}, args...)...)
return err
}
require.Error(t, run("stop"), "stop without a running profile must fail")
require.NoError(t, run("start"))
assert.Error(t, run("start"), "second start must be rejected while profiling")
require.NoError(t, run("stop"))
assert.Error(t, run("stop"), "second stop must be rejected")
assert.NoError(t, run("start"), "profiling can be started again after a stop")
assert.NoError(t, run("stop"))
}
// TestDebugForKeepsRunningCPUProfile covers `debug for` started while a
// profile from `debug cpu start` is running: it must say so, leave the
// profile alone, and still create the bundle.
func TestDebugForKeepsRunningCPUProfile(t *testing.T) {
addr := startDebugTestDaemon(t)
_, err := runDebugCmd(addr, "cpu", "start")
require.NoError(t, err)
out, err := runDebugCmd(addr, "for", "1s", "-S=false", "--no-updown")
require.NoError(t, err, "output: %s", out)
assert.Contains(t, out, "CPU profiling is already running", "the conflict must be explained")
assert.NotContains(t, out, "rpc error", "the raw RPC error must not reach the user")
assert.Contains(t, out, "Local file:", "the bundle must still be created")
_, err = runDebugCmd(addr, "cpu", "stop")
assert.NoError(t, err, "the profile started by the user must still be running")
}
func TestDebugForNoUpDown(t *testing.T) {
addr := startDebugTestDaemon(t)
out, err := runDebugCmd(addr, "for", "1s", "-S=false", "--no-updown")
require.NoError(t, err, "output: %s", out)
assert.NotContains(t, out, "netbird down", "--no-updown must not bring the daemon down")
assert.NotContains(t, out, "netbird up", "--no-updown must not bring the daemon up")
assert.Contains(t, out, "Local file:", "the bundle must still be created")
}
+50
View File
@@ -7,6 +7,7 @@ import (
"fmt" "fmt"
"net/http" "net/http"
"runtime" "runtime"
"slices"
"strings" "strings"
"sync" "sync"
@@ -25,6 +26,30 @@ var serviceCmd = &cobra.Command{
const defaultJSONSocket = "unix:///var/run/netbird-http.sock" const defaultJSONSocket = "unix:///var/run/netbird-http.sock"
// forbiddenServiceEnvVars are the environment variables the service is never
// registered with, keyed in upper case since these are Windows names. Each one
// decides where the daemon resolves something it then uses with the privileges
// of the account it runs under — LocalSystem on Windows, root elsewhere: the
// executables it runs (PATH, PATHEXT, COMSPEC, SystemRoot, windir) or the
// directory it writes temporary files in (TEMP, TMP). The daemon needs none of
// them, and the utilities it shells out to are resolved by absolute path.
var forbiddenServiceEnvVars = map[string]struct{}{
"PATH": {},
"PATHEXT": {},
"SYSTEMROOT": {},
"WINDIR": {},
"COMSPEC": {},
"TEMP": {},
"TMP": {},
}
// forbiddenServiceEnvPrefixes are the dynamic-loader families, refused whole
// rather than by name: LD_PRELOAD, DYLD_INSERT_LIBRARIES and their siblings all
// reach the loader of the process, the set differs per platform and libc, and
// new members arrive with new OS releases. Listing them one by one is a list
// that is wrong the moment it is written.
var forbiddenServiceEnvPrefixes = []string{"LD_", "DYLD_"}
var ( var (
serviceName string serviceName string
serviceEnvVars []string serviceEnvVars []string
@@ -127,8 +152,33 @@ func parseServiceEnvVars(envVars []string) (map[string]string, error) {
return nil, fmt.Errorf("empty environment variable key in: %s", env) return nil, fmt.Errorf("empty environment variable key in: %s", env)
} }
if isForbiddenServiceEnvVar(key) {
return nil, fmt.Errorf("environment variable %s cannot be set on the service: it decides where the service resolves the executables, libraries or temporary files it uses", key)
}
envMap[key] = value envMap[key] = value
} }
return envMap, nil return envMap, nil
} }
// isForbiddenServiceEnvVar reports whether name is one the service must not be
// registered with.
//
// The names are matched case-insensitively only on Windows, where they are the
// same variable however they are spelled. Elsewhere the environment is
// case-sensitive, so Path and PATH are two different variables and only the
// exact spelling is the one the loader reads.
func isForbiddenServiceEnvVar(name string) bool {
if runtime.GOOS == "windows" {
name = strings.ToUpper(name)
}
if _, forbidden := forbiddenServiceEnvVars[name]; forbidden {
return true
}
return slices.ContainsFunc(forbiddenServiceEnvPrefixes, func(prefix string) bool {
return strings.HasPrefix(name, prefix)
})
}
+50 -6
View File
@@ -14,6 +14,7 @@ import (
"github.com/netbirdio/netbird/client/configs" "github.com/netbirdio/netbird/client/configs"
"github.com/netbirdio/netbird/client/internal/daemonaddr" "github.com/netbirdio/netbird/client/internal/daemonaddr"
"github.com/netbirdio/netbird/client/internal/elevate"
"github.com/netbirdio/netbird/util" "github.com/netbirdio/netbird/util"
) )
@@ -43,10 +44,33 @@ func serviceParamsPath() string {
// loadServiceParams reads saved service parameters from disk. // loadServiceParams reads saved service parameters from disk.
// Returns nil with no error if the file does not exist. // Returns nil with no error if the file does not exist.
//
// The file is read by an elevated install and decides the arguments and the
// environment of the service it then registers, so it is used only when its
// ownership and permissions are the ones saveServiceParams leaves behind. That
// restricted ACL is applied when the file is written, which is not necessarily
// before it is first read, so this is checked rather than assumed. A file that
// fails the check is treated as absent, and the install proceeds with its
// defaults.
func loadServiceParams() (*serviceParams, error) { func loadServiceParams() (*serviceParams, error) {
path := serviceParamsPath() path := serviceParamsPath()
data, err := os.ReadFile(path) // Resolve links first so the checks apply to the file that is actually read.
// Since the check covers every directory above it as well, nobody who fails
// it can swap the file between here and the read below.
resolved, err := filepath.EvalSymlinks(path)
if err != nil {
if os.IsNotExist(err) {
return nil, nil //nolint:nilnil
}
return nil, fmt.Errorf("resolve service params %s: %w", path, err)
}
if err := elevate.CheckOnlyOwnerWritable(resolved); err != nil {
return nil, fmt.Errorf("refusing to read service params from %s: %w", resolved, err)
}
data, err := os.ReadFile(resolved)
if err != nil { if err != nil {
if os.IsNotExist(err) { if os.IsNotExist(err) {
return nil, nil //nolint:nilnil return nil, nil //nolint:nilnil
@@ -182,10 +206,16 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) {
// If --service-env was explicitly set to empty, all saved env vars are cleared. // If --service-env was explicitly set to empty, all saved env vars are cleared.
// If --service-env was not set, saved env vars are used entirely. // If --service-env was not set, saved env vars are used entirely.
func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) { func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
// A forbidden name explicitly passed on the command line is an error the
// operator is told about, but one restored from a file written by an older
// version is dropped: an install that refuses to run would leave the host
// without a daemon over a variable nobody is asking for any more.
saved := dropForbiddenServiceEnvVars(cmd, params.ServiceEnvVars)
if !cmd.Flags().Changed("service-env") { if !cmd.Flags().Changed("service-env") {
if len(params.ServiceEnvVars) > 0 { if len(saved) > 0 {
// No explicit env vars: rebuild serviceEnvVars from saved params. // No explicit env vars: rebuild serviceEnvVars from saved params.
serviceEnvVars = envMapToSlice(params.ServiceEnvVars) serviceEnvVars = envMapToSlice(saved)
} }
return return
} }
@@ -204,13 +234,13 @@ func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
return return
} }
if len(params.ServiceEnvVars) == 0 { if len(saved) == 0 {
return return
} }
// Merge saved values underneath explicit ones. // Merge saved values underneath explicit ones.
merged := make(map[string]string, len(params.ServiceEnvVars)+len(explicit)) merged := make(map[string]string, len(saved)+len(explicit))
maps.Copy(merged, params.ServiceEnvVars) maps.Copy(merged, saved)
maps.Copy(merged, explicit) // explicit wins on conflict maps.Copy(merged, explicit) // explicit wins on conflict
serviceEnvVars = envMapToSlice(merged) serviceEnvVars = envMapToSlice(merged)
} }
@@ -233,6 +263,20 @@ var resetParamsCmd = &cobra.Command{
}, },
} }
// dropForbiddenServiceEnvVars returns the saved entries that may still be
// registered on the service, reporting every one it leaves behind.
func dropForbiddenServiceEnvVars(cmd *cobra.Command, saved map[string]string) map[string]string {
kept := make(map[string]string, len(saved))
for key, value := range saved {
if isForbiddenServiceEnvVar(key) {
cmd.PrintErrf("Warning: ignoring saved service environment variable %s: it decides where the service resolves the executables, libraries or temporary files it uses\n", key)
continue
}
kept[key] = value
}
return kept
}
// envMapToSlice converts a map of env vars to a KEY=VALUE slice. // envMapToSlice converts a map of env vars to a KEY=VALUE slice.
func envMapToSlice(m map[string]string) []string { func envMapToSlice(m map[string]string) []string {
s := make([]string, 0, len(m)) s := make([]string, 0, len(m))
+54
View File
@@ -9,6 +9,7 @@ import (
"go/token" "go/token"
"os" "os"
"path/filepath" "path/filepath"
"runtime"
"strings" "strings"
"testing" "testing"
@@ -353,6 +354,59 @@ func TestApplyServiceEnvParams_NotChanged(t *testing.T) {
assert.Equal(t, map[string]string{"FROM_SAVED": "val"}, result) assert.Equal(t, map[string]string{"FROM_SAVED": "val"}, result)
} }
func TestParseServiceEnvVars_RejectsForbiddenNames(t *testing.T) {
for _, env := range []string{"PATH=C:\\somewhere", "LD_PRELOAD=/tmp/lib.so", "DYLD_FALLBACK_LIBRARY_PATH=/tmp"} {
_, err := parseServiceEnvVars([]string{"KEEP=me", env})
require.Errorf(t, err, "%s selects what the service resolves and must be refused", env)
}
}
func TestIsForbiddenServiceEnvVar(t *testing.T) {
// The loader families are matched by prefix, so a name nobody has heard of
// yet is refused too.
for _, name := range []string{
"PATH", "PATHEXT", "COMSPEC", "SYSTEMROOT", "WINDIR", "TEMP", "TMP",
"LD_PRELOAD", "LD_AUDIT", "DYLD_INSERT_LIBRARIES", "DYLD_FALLBACK_FRAMEWORK_PATH",
} {
assert.Truef(t, isForbiddenServiceEnvVar(name), "%s must be refused", name)
}
// The prefix must not swallow names that merely start with the same letters.
for _, name := range []string{"NB_LOG_LEVEL", "NB_WG_DEBUG", "HTTPS_PROXY", "LDAP_URL", "DYLDX"} {
assert.Falsef(t, isForbiddenServiceEnvVar(name), "%s has no reason to be refused", name)
}
// On Windows a variable is the same one however it is spelled; elsewhere
// Path and PATH are two variables and only the exact one is read.
if runtime.GOOS == "windows" {
assert.True(t, isForbiddenServiceEnvVar("Path"))
assert.True(t, isForbiddenServiceEnvVar("ld_preload"))
} else {
assert.False(t, isForbiddenServiceEnvVar("Path"))
assert.False(t, isForbiddenServiceEnvVar("ld_preload"))
}
}
func TestApplyServiceEnvParams_DropsForbiddenSavedNames(t *testing.T) {
origServiceEnvVars := serviceEnvVars
t.Cleanup(func() { serviceEnvVars = origServiceEnvVars })
serviceEnvVars = nil
cmd := &cobra.Command{}
cmd.Flags().StringSlice("service-env", nil, "")
saved := &serviceParams{
ServiceEnvVars: map[string]string{"PATH": "C:\\attacker", "NB_LOG_FORMAT": "json"},
}
applyServiceEnvParams(cmd, saved)
result, err := parseServiceEnvVars(serviceEnvVars)
require.NoError(t, err, "a saved PATH must be dropped rather than fail the install")
assert.Equal(t, map[string]string{"NB_LOG_FORMAT": "json"}, result)
}
func TestApplyServiceEnvParams_ExplicitEmptyClears(t *testing.T) { func TestApplyServiceEnvParams_ExplicitEmptyClears(t *testing.T) {
origServiceEnvVars := serviceEnvVars origServiceEnvVars := serviceEnvVars
t.Cleanup(func() { serviceEnvVars = origServiceEnvVars }) t.Cleanup(func() { serviceEnvVars = origServiceEnvVars })
+57
View File
@@ -0,0 +1,57 @@
//go:build !windows && !ios && !android
package cmd
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/configs"
)
// The Windows equivalent of this is the ACL check in
// elevate.CheckOnlyOwnerWritable, covered by that package's own tests; here the
// point is that loadServiceParams asks the question at all.
func TestLoadServiceParams_RefusesWorldWritableFile(t *testing.T) {
tmpDir := t.TempDir()
original := configs.StateDir
t.Cleanup(func() { configs.StateDir = original })
configs.StateDir = tmpDir
path := filepath.Join(tmpDir, serviceParamsFile)
require.NoError(t, os.WriteFile(path, []byte(`{"log_level":"debug"}`), 0o666))
// WriteFile is subject to the umask, so set the bits that matter explicitly.
require.NoError(t, os.Chmod(path, 0o666))
params, err := loadServiceParams()
require.Error(t, err, "a service.json anyone can rewrite must not be trusted")
assert.Nil(t, params)
require.NoError(t, os.Chmod(path, 0o600))
params, err = loadServiceParams()
require.NoError(t, err)
require.NotNil(t, params)
assert.Equal(t, "debug", params.LogLevel)
}
func TestLoadServiceParams_RefusesWorldWritableDirectory(t *testing.T) {
tmpDir := t.TempDir()
stateDir := filepath.Join(tmpDir, "state")
require.NoError(t, os.Mkdir(stateDir, 0o777))
require.NoError(t, os.Chmod(stateDir, 0o777))
original := configs.StateDir
t.Cleanup(func() { configs.StateDir = original })
configs.StateDir = stateDir
require.NoError(t, os.WriteFile(filepath.Join(stateDir, serviceParamsFile), []byte(`{}`), 0o600))
params, err := loadServiceParams()
require.Error(t, err, "a service.json in a directory anyone can replace entries in must not be trusted")
assert.Nil(t, params)
}
+2
View File
@@ -24,6 +24,7 @@ const (
tableFilter = "filter" tableFilter = "filter"
tableNat = "nat" tableNat = "nat"
tableMangle = "mangle" tableMangle = "mangle"
tableRaw = "raw"
// chainACLInput is the peer ACL chain that holds installed // chainACLInput is the peer ACL chain that holds installed
// peer-filtering rules. // peer-filtering rules.
@@ -34,6 +35,7 @@ const (
mangleForwardKey chainKey = "MANGLE-FORWARD" mangleForwardKey chainKey = "MANGLE-FORWARD"
chainInput = "INPUT" chainInput = "INPUT"
chainOutput = "OUTPUT"
chainPostrouting = "POSTROUTING" chainPostrouting = "POSTROUTING"
chainPrerouting = "PREROUTING" chainPrerouting = "PREROUTING"
chainForward = "FORWARD" chainForward = "FORWARD"
+2 -139
View File
@@ -25,9 +25,8 @@ type Manager struct {
wgIface iFaceMapper wgIface iFaceMapper
ipv4Client *iptables.IPTables ipv4Client *iptables.IPTables
family4 *family family4 *family
rawSupported bool
// IPv6 counterparts, nil when no v6 overlay // IPv6 counterparts, nil when no v6 overlay
ipv6Client *iptables.IPTables ipv6Client *iptables.IPTables
@@ -108,10 +107,6 @@ func (m *Manager) Init(stateManager *statemanager.Manager) error {
return err return err
} }
if err := m.initNoTrackChain(); err != nil {
log.Warnf("raw table not available, notrack rules will be disabled: %v", err)
}
// Trust after all fatal init steps so a later failure doesn't leave the // Trust after all fatal init steps so a later failure doesn't leave the
// interface in firewalld's trusted zone without a corresponding Close. // interface in firewalld's trusted zone without a corresponding Close.
if err := firewalld.TrustInterface(m.wgIface.Name()); err != nil { if err := firewalld.TrustInterface(m.wgIface.Name()); err != nil {
@@ -285,10 +280,6 @@ func (m *Manager) Close(stateManager *statemanager.Manager) error {
var merr *multierror.Error var merr *multierror.Error
if err := m.cleanupNoTrackChain(); err != nil {
merr = multierror.Append(merr, fmt.Errorf("cleanup notrack chain: %w", err))
}
if m.hasIPv6() { if m.hasIPv6() {
if err := m.family6.Reset(); err != nil { if err := m.family6.Reset(); err != nil {
merr = multierror.Append(merr, fmt.Errorf("reset v6 family: %w", err)) merr = multierror.Append(merr, fmt.Errorf("reset v6 family: %w", err))
@@ -440,134 +431,6 @@ func (m *Manager) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Proto
return m.family4.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort) return m.family4.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort)
} }
const (
chainNameRaw = "NETBIRD-RAW"
chainOutput = "OUTPUT"
tableRaw = "raw"
)
// SetupEBPFProxyNoTrack creates notrack rules for eBPF proxy loopback traffic.
// This prevents conntrack from tracking WireGuard proxy traffic on loopback, which
// can interfere with MASQUERADE rules (e.g., from container runtimes like Podman/netavark).
//
// Traffic flows that need NOTRACK:
//
// 1. Egress: WireGuard -> fake endpoint (before eBPF rewrite)
// src=127.0.0.1:wgPort -> dst=127.0.0.1:fakePort
// Matched by: sport=wgPort
//
// 2. Egress: Proxy -> WireGuard (via raw socket)
// src=127.0.0.1:fakePort -> dst=127.0.0.1:wgPort
// Matched by: dport=wgPort
//
// 3. Ingress: Packets to WireGuard
// dst=127.0.0.1:wgPort
// Matched by: dport=wgPort
//
// 4. Ingress: Packets to proxy (after eBPF rewrite)
// dst=127.0.0.1:proxyPort
// Matched by: dport=proxyPort
//
// Rules are cleaned up when the firewall manager is closed.
func (m *Manager) SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error {
m.mutex.Lock()
defer m.mutex.Unlock()
if !m.rawSupported {
return fmt.Errorf("raw table not available")
}
wgPortStr := fmt.Sprintf("%d", wgPort)
proxyPortStr := fmt.Sprintf("%d", proxyPort)
// Egress rules: match outgoing loopback UDP packets
outputRuleSport := []string{"-o", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--sport", wgPortStr, "-j", "NOTRACK"}
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, outputRuleSport...); err != nil {
return fmt.Errorf("add output sport notrack rule: %w", err)
}
outputRuleDport := []string{"-o", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", wgPortStr, "-j", "NOTRACK"}
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, outputRuleDport...); err != nil {
return fmt.Errorf("add output dport notrack rule: %w", err)
}
// Ingress rules: match incoming loopback UDP packets
preroutingRuleWg := []string{"-i", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", wgPortStr, "-j", "NOTRACK"}
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, preroutingRuleWg...); err != nil {
return fmt.Errorf("add prerouting wg notrack rule: %w", err)
}
preroutingRuleProxy := []string{"-i", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", proxyPortStr, "-j", "NOTRACK"}
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, preroutingRuleProxy...); err != nil {
return fmt.Errorf("add prerouting proxy notrack rule: %w", err)
}
log.Debugf("set up ebpf proxy notrack rules for ports %d,%d", proxyPort, wgPort)
return nil
}
func (m *Manager) initNoTrackChain() error {
if err := m.cleanupNoTrackChain(); err != nil {
log.Debugf("cleanup notrack chain: %v", err)
}
if err := m.ipv4Client.NewChain(tableRaw, chainNameRaw); err != nil {
return fmt.Errorf("create chain: %w", err)
}
jumpRule := []string{"-j", chainNameRaw}
if err := m.ipv4Client.InsertUnique(tableRaw, chainOutput, 1, jumpRule...); err != nil {
if delErr := m.ipv4Client.DeleteChain(tableRaw, chainNameRaw); delErr != nil {
log.Debugf("delete orphan chain: %v", delErr)
}
return fmt.Errorf("add output jump rule: %w", err)
}
if err := m.ipv4Client.InsertUnique(tableRaw, chainPrerouting, 1, jumpRule...); err != nil {
if delErr := m.ipv4Client.DeleteIfExists(tableRaw, chainOutput, jumpRule...); delErr != nil {
log.Debugf("delete output jump rule: %v", delErr)
}
if delErr := m.ipv4Client.DeleteChain(tableRaw, chainNameRaw); delErr != nil {
log.Debugf("delete orphan chain: %v", delErr)
}
return fmt.Errorf("add prerouting jump rule: %w", err)
}
m.rawSupported = true
return nil
}
func (m *Manager) cleanupNoTrackChain() error {
exists, err := m.ipv4Client.ChainExists(tableRaw, chainNameRaw)
if err != nil {
if !m.rawSupported {
return nil
}
return fmt.Errorf("check chain exists: %w", err)
}
if !exists {
return nil
}
jumpRule := []string{"-j", chainNameRaw}
if err := m.ipv4Client.DeleteIfExists(tableRaw, chainOutput, jumpRule...); err != nil {
return fmt.Errorf("remove output jump rule: %w", err)
}
if err := m.ipv4Client.DeleteIfExists(tableRaw, chainPrerouting, jumpRule...); err != nil {
return fmt.Errorf("remove prerouting jump rule: %w", err)
}
if err := m.ipv4Client.ClearAndDeleteChain(tableRaw, chainNameRaw); err != nil {
return fmt.Errorf("clear and delete chain: %w", err)
}
m.rawSupported = false
return nil
}
func getConntrackEstablished() []string { func getConntrackEstablished() []string {
return []string{"-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"} return []string{"-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"}
} }
-4
View File
@@ -192,10 +192,6 @@ type Manager interface {
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule. // RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
RemoveOutputDNAT(localAddr netip.Addr, protocol Protocol, originalPort, translatedPort uint16) error RemoveOutputDNAT(localAddr netip.Addr, protocol Protocol, originalPort, translatedPort uint16) error
// SetupEBPFProxyNoTrack creates static notrack rules for eBPF proxy loopback traffic.
// This prevents conntrack from interfering with WireGuard proxy communication.
SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error
} }
// GenKey builds the rule id for this pair from the given format. // GenKey builds the rule id for this pair from the given format.
-182
View File
@@ -12,7 +12,6 @@ import (
"github.com/google/nftables/expr" "github.com/google/nftables/expr"
"github.com/hashicorp/go-multierror" "github.com/hashicorp/go-multierror"
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
"golang.org/x/sys/unix"
nberrors "github.com/netbirdio/netbird/client/errors" nberrors "github.com/netbirdio/netbird/client/errors"
firewall "github.com/netbirdio/netbird/client/firewall/manager" firewall "github.com/netbirdio/netbird/client/firewall/manager"
@@ -55,9 +54,6 @@ type Manager struct {
// IPv6 counterpart, nil when no v6 overlay. // IPv6 counterpart, nil when no v6 overlay.
family6 *family family6 *family
notrackOutputChain *nftables.Chain
notrackPreroutingChain *nftables.Chain
extMonitor *externalChainMonitor extMonitor *externalChainMonitor
} }
@@ -170,10 +166,6 @@ func (m *Manager) initFirewall() (err error) {
} }
} }
if err := m.initNoTrackChains(workTable); err != nil {
log.Warnf("raw priority chains not available, notrack rules will be disabled: %v", err)
}
return nil return nil
} }
@@ -455,10 +447,6 @@ func (m *Manager) Flush() error {
} }
} }
if err := m.refreshNoTrackChains(); err != nil {
log.Errorf("failed to refresh notrack chains: %v", err)
}
return nil return nil
} }
@@ -571,176 +559,6 @@ func (m *Manager) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Proto
return m.family4.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort) return m.family4.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort)
} }
const (
chainNameRawOutput = "netbird-raw-out"
chainNameRawPrerouting = "netbird-raw-pre"
)
// SetupEBPFProxyNoTrack creates notrack rules for eBPF proxy loopback traffic.
// This prevents conntrack from tracking WireGuard proxy traffic on loopback, which
// can interfere with MASQUERADE rules (e.g., from container runtimes like Podman/netavark).
//
// Traffic flows that need NOTRACK:
//
// 1. Egress: WireGuard -> fake endpoint (before eBPF rewrite)
// src=127.0.0.1:wgPort -> dst=127.0.0.1:fakePort
// Matched by: sport=wgPort
//
// 2. Egress: Proxy -> WireGuard (via raw socket)
// src=127.0.0.1:fakePort -> dst=127.0.0.1:wgPort
// Matched by: dport=wgPort
//
// 3. Ingress: Packets to WireGuard
// dst=127.0.0.1:wgPort
// Matched by: dport=wgPort
//
// 4. Ingress: Packets to proxy (after eBPF rewrite)
// dst=127.0.0.1:proxyPort
// Matched by: dport=proxyPort
//
// Rules are cleaned up when the firewall manager is closed.
func (m *Manager) SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error {
m.mutex.Lock()
defer m.mutex.Unlock()
if m.notrackOutputChain == nil || m.notrackPreroutingChain == nil {
return fmt.Errorf("notrack chains not initialized")
}
proxyPortBytes := binaryutil.BigEndian.PutUint16(proxyPort)
wgPortBytes := binaryutil.BigEndian.PutUint16(wgPort)
loopback := []byte{127, 0, 0, 1}
// Egress rules: match outgoing loopback UDP packets
m.rConn.AddRule(&nftables.Rule{
Table: m.notrackOutputChain.Table,
Chain: m.notrackOutputChain,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // sport=wgPort
&expr.Counter{},
&expr.Notrack{},
},
})
m.rConn.AddRule(&nftables.Rule{
Table: m.notrackOutputChain.Table,
Chain: m.notrackOutputChain,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // dport=wgPort
&expr.Counter{},
&expr.Notrack{},
},
})
// Ingress rules: match incoming loopback UDP packets
m.rConn.AddRule(&nftables.Rule{
Table: m.notrackPreroutingChain.Table,
Chain: m.notrackPreroutingChain,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // dport=wgPort
&expr.Counter{},
&expr.Notrack{},
},
})
m.rConn.AddRule(&nftables.Rule{
Table: m.notrackPreroutingChain.Table,
Chain: m.notrackPreroutingChain,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: proxyPortBytes}, // dport=proxyPort
&expr.Counter{},
&expr.Notrack{},
},
})
if err := m.rConn.Flush(); err != nil {
return fmt.Errorf("flush notrack rules: %w", err)
}
log.Debugf("set up ebpf proxy notrack rules for ports %d,%d", proxyPort, wgPort)
return nil
}
func (m *Manager) initNoTrackChains(table *nftables.Table) error {
m.notrackOutputChain = m.rConn.AddChain(&nftables.Chain{
Name: chainNameRawOutput,
Table: table,
Type: nftables.ChainTypeFilter,
Hooknum: nftables.ChainHookOutput,
Priority: nftables.ChainPriorityRaw,
})
m.notrackPreroutingChain = m.rConn.AddChain(&nftables.Chain{
Name: chainNameRawPrerouting,
Table: table,
Type: nftables.ChainTypeFilter,
Hooknum: nftables.ChainHookPrerouting,
Priority: nftables.ChainPriorityRaw,
})
if err := m.rConn.Flush(); err != nil {
return fmt.Errorf("flush chain creation: %w", err)
}
return nil
}
func (m *Manager) refreshNoTrackChains() error {
chains, err := m.rConn.ListChainsOfTableFamily(nftables.TableFamilyIPv4)
if err != nil {
return fmt.Errorf("list chains: %w", err)
}
tableName := getTableName()
for _, c := range chains {
if c.Table.Name != tableName {
continue
}
switch c.Name {
case chainNameRawOutput:
m.notrackOutputChain = c
case chainNameRawPrerouting:
m.notrackPreroutingChain = c
}
}
return nil
}
func (m *Manager) createWorkTable() (*nftables.Table, error) { func (m *Manager) createWorkTable() (*nftables.Table, error) {
return m.createWorkTableFamily(nftables.TableFamilyIPv4) return m.createWorkTableFamily(nftables.TableFamilyIPv4)
} }
+1 -1
View File
@@ -192,7 +192,7 @@ func (r *family) addPostroutingRules() {
Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkMasquerade), Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkMasquerade),
}, },
// We need to exclude the loopback interface as this changes the ebpf proxy port // We need to exclude the loopback interface as this changes the wg proxy port
&expr.Meta{ &expr.Meta{
Key: expr.MetaKeyOIFNAME, Key: expr.MetaKeyOIFNAME,
Register: 1, Register: 1,
-6
View File
@@ -879,12 +879,6 @@ func (m *Manager) resetState() {
} }
} }
// SetupEBPFProxyNoTrack is not supported by the userspace firewall: eBPF isn't
// used in userspace mode, so this should never be called.
func (m *Manager) SetupEBPFProxyNoTrack(uint16, uint16) error {
return errNotSupported
}
// UpdateSet updates the rule destinations associated with the given set // UpdateSet updates the rule destinations associated with the given set
// by merging the existing prefixes with the new ones, then deduplicating. // by merging the existing prefixes with the new ones, then deduplicating.
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error { func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
@@ -9,6 +9,7 @@ import (
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
nberrors "github.com/netbirdio/netbird/client/errors" nberrors "github.com/netbirdio/netbird/client/errors"
"github.com/netbirdio/netbird/client/internal/wincmd"
) )
type action string type action string
@@ -91,7 +92,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err
if action == addRule { if action == addRule {
args = append(args, extraArgs...) args = append(args, extraArgs...)
} }
netshCmd := GetSystem32Command("netsh") netshCmd := wincmd.System32("netsh")
cmd := exec.Command(netshCmd, args...) cmd := exec.Command(netshCmd, args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true} cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
return cmd.Run() return cmd.Run()
@@ -100,7 +101,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err
func isWindowsFirewallReachable() bool { func isWindowsFirewallReachable() bool {
args := []string{"advfirewall", "show", "allprofiles", "state"} args := []string{"advfirewall", "show", "allprofiles", "state"}
netshCmd := GetSystem32Command("netsh") netshCmd := wincmd.System32("netsh")
cmd := exec.Command(netshCmd, args...) cmd := exec.Command(netshCmd, args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true} cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
@@ -117,23 +118,10 @@ func isWindowsFirewallReachable() bool {
func isFirewallRuleActive(ruleName string) bool { func isFirewallRuleActive(ruleName string) bool {
args := []string{"advfirewall", "firewall", "show", "rule", "name=" + ruleName} args := []string{"advfirewall", "firewall", "show", "rule", "name=" + ruleName}
netshCmd := GetSystem32Command("netsh") netshCmd := wincmd.System32("netsh")
cmd := exec.Command(netshCmd, args...) cmd := exec.Command(netshCmd, args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true} cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
_, err := cmd.Output() _, err := cmd.Output()
return err == nil return err == nil
} }
// GetSystem32Command checks if a command can be found in the system path and returns it. In case it can't find it
// in the path it will return the full path of a command assuming C:\windows\system32 as the base path.
func GetSystem32Command(command string) string {
_, err := exec.LookPath(command)
if err == nil {
return command
}
log.Tracef("Command %s not found in PATH, using C:\\windows\\system32\\%s.exe path", command, command)
return "C:\\windows\\system32\\" + command + ".exe"
}
+226
View File
@@ -0,0 +1,226 @@
package configurer
import (
"net"
"net/netip"
"slices"
"sync"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
// allowedIPStore mirrors the allowed IPs configured on each peer of a device.
//
// A configurer is the only writer of its device's peer set, so the mirror is authoritative
// by construction. It spares the paths that have to rewrite one peer's allowed IPs a full
// device dump just to recover prefixes the process already configured itself.
//
// An allowed IP belongs to exactly one peer: configuring a prefix on a peer takes it away
// from whichever peer held it before, and the configurer leaves that handover to the device
// rather than removing the prefix from the previous holder itself. The store tracks the
// owner of each prefix and performs the same handover, so rewriting one peer's list never
// takes a prefix back from the peer that owns it now.
//
// Its own lock guards the map alone, not the device write it accompanies. Consistency
// between the two rests on the caller serializing every configurer call, which WGIface
// does with its mutex; two unserialized writers would interleave a device write with the
// record of a different one.
//
// An operator reconfiguring the device out of band, through `wg set` or the UAPI socket,
// is the one way the mirror can still go stale. A peer missing from it falls back to the
// device, which reseats that peer's prefixes and their ownership; a peer that is present
// does not, so one recorded from empty while the device already held prefixes keeps only
// what was recorded, and the next endpoint removal drops the rest.
type allowedIPStore struct {
mu sync.RWMutex
peers map[wgtypes.Key][]netip.Prefix
owners map[netip.Prefix]wgtypes.Key
}
func newAllowedIPStore() *allowedIPStore {
return &allowedIPStore{
peers: make(map[wgtypes.Key][]netip.Prefix),
owners: make(map[netip.Prefix]wgtypes.Key),
}
}
// get returns the prefixes recorded for a peer, and whether the peer is known at all.
// The caller receives a copy and may retain or modify it freely.
func (s *allowedIPStore) get(key wgtypes.Key) ([]netip.Prefix, bool) {
s.mu.RLock()
defer s.mu.RUnlock()
prefixes, ok := s.peers[key]
if !ok {
return nil, false
}
return slices.Clone(prefixes), true
}
// set replaces the prefixes recorded for a peer.
func (s *allowedIPStore) set(key wgtypes.Key, prefixes []netip.Prefix) {
s.mu.Lock()
defer s.mu.Unlock()
k := key
s.releaseLocked(k)
normalized := normalizePrefixes(prefixes)
for _, prefix := range normalized {
s.claimLocked(k, prefix)
}
s.peers[k] = normalized
}
// add records prefixes on a peer without dropping the ones already there, matching the
// union semantics of a peer update that does not replace its allowed IPs. It records the
// peer if it is not known yet, so it belongs to the operations that create a peer on the
// device rather than to the update-only ones.
func (s *allowedIPStore) add(key wgtypes.Key, prefixes []netip.Prefix) {
s.mu.Lock()
defer s.mu.Unlock()
s.mergeLocked(key, prefixes)
}
// addExisting is add for an update-only device operation. Such an operation is a silent
// no-op when the peer is absent, so recording a peer here would leave the store claiming
// prefixes the device never took, and the peer would then be recreated by the next endpoint
// removal, stealing those allowed IPs from the peer that legitimately holds them.
func (s *allowedIPStore) addExisting(key wgtypes.Key, prefixes []netip.Prefix) {
s.mu.Lock()
defer s.mu.Unlock()
k := key
if _, ok := s.peers[k]; !ok {
return
}
s.mergeLocked(k, prefixes)
}
// ensure records a peer with no prefixes unless it is already known. A device operation
// that is not update-only creates the peer when it is absent, so it has to be recorded even
// when it configures nothing else; otherwise the peer exists on the device while the store
// treats it as unknown, and a prefix later handed over to it is not accounted for.
func (s *allowedIPStore) ensure(key wgtypes.Key) {
s.mu.Lock()
defer s.mu.Unlock()
k := key
if _, ok := s.peers[k]; !ok {
s.peers[k] = nil
}
}
// forget drops every prefix recorded for a peer.
func (s *allowedIPStore) forget(key wgtypes.Key) {
s.mu.Lock()
defer s.mu.Unlock()
k := key
s.releaseLocked(k)
delete(s.peers, k)
}
// reset drops every peer, mirroring a device reconfiguration that replaces the peer set.
func (s *allowedIPStore) reset() {
s.mu.Lock()
defer s.mu.Unlock()
s.peers = make(map[wgtypes.Key][]netip.Prefix)
s.owners = make(map[netip.Prefix]wgtypes.Key)
}
// mergeLocked unions normalized prefixes into a peer and transfers their ownership.
// The caller must hold s.mu for writing.
func (s *allowedIPStore) mergeLocked(k wgtypes.Key, prefixes []netip.Prefix) {
merged := s.peers[k]
for _, prefix := range prefixes {
prefix = normalizePrefix(prefix)
s.claimLocked(k, prefix)
if !slices.Contains(merged, prefix) {
merged = append(merged, prefix)
}
}
s.peers[k] = merged
}
// claimLocked hands a prefix over to a peer, taking it from its previous owner the way the
// device does when the same prefix is configured on a second peer.
func (s *allowedIPStore) claimLocked(k wgtypes.Key, prefix netip.Prefix) {
if owner, ok := s.owners[prefix]; ok && owner != k {
s.peers[owner] = slices.DeleteFunc(s.peers[owner], func(p netip.Prefix) bool {
return p == prefix
})
}
s.owners[prefix] = k
}
// releaseLocked drops a peer's claim on every prefix it currently holds.
func (s *allowedIPStore) releaseLocked(k wgtypes.Key) {
for _, prefix := range s.peers[k] {
if s.owners[prefix] == k {
delete(s.owners, prefix)
}
}
}
// normalizePrefix puts a prefix into the form the store recognises it by. It clears the
// host bits, which a device does on its own, so a caller passing 10.20.0.1/16 still matches
// the 10.20.0.0/16 read back from the device; and it unmaps a v4-mapped prefix so that it
// compares equal to, and marshals like, the plain v4 prefix for the same network.
//
// Masking comes first because it also decides the address family: only a prefix at least 96
// bits long keeps the mapped marker through the mask, so a shorter prefix inside the mapped
// range is a genuine v6 prefix and unmapping it would yield an invalid v4 prefix.
func normalizePrefix(prefix netip.Prefix) netip.Prefix {
masked := prefix.Masked()
addr := masked.Addr()
if !addr.Is4In6() {
return masked
}
return netip.PrefixFrom(addr.Unmap(), masked.Bits()-96)
}
// normalizePrefixes returns a normalized copy without changing the caller's slice.
func normalizePrefixes(prefixes []netip.Prefix) []netip.Prefix {
normalized := make([]netip.Prefix, len(prefixes))
for i, prefix := range prefixes {
normalized[i] = normalizePrefix(prefix)
}
return normalized
}
// ipNetsToPrefixes converts addresses read back from a device. Unmap keeps a v4-mapped v6
// address comparable to the plain v4 prefix the configurer was given.
func ipNetsToPrefixes(ipNets []net.IPNet) []netip.Prefix {
prefixes := make([]netip.Prefix, 0, len(ipNets))
for _, ipNet := range ipNets {
addr, ok := netip.AddrFromSlice(ipNet.IP)
if !ok {
continue
}
ones, maskBits := ipNet.Mask.Size()
// A device may report a v4 prefix as a v4-mapped address. Align the address form with
// the mask rather than unmapping on sight: a 32 bit mask always describes v4, while a
// 128 bit mask describes v4 only when it covers the mapped prefix, so a genuine v6
// prefix inside the mapped range stays v6 instead of being dropped as invalid.
if addr.Is4In6() {
switch {
case maskBits == 32:
addr = addr.Unmap()
case maskBits == 128 && ones >= 96:
addr, ones = addr.Unmap(), ones-96
}
}
prefix := netip.PrefixFrom(addr, ones)
if !prefix.IsValid() {
continue
}
prefixes = append(prefixes, prefix.Masked())
}
return prefixes
}
+263
View File
@@ -0,0 +1,263 @@
package configurer
import (
"net"
"net/netip"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
// The store keys on the parsed key, so the tests use two distinct ones rather than names.
var (
testPeer = wgtypes.Key{1}
otherPeer = wgtypes.Key{2}
)
func TestAllowedIPStoreUnknownPeer(t *testing.T) {
s := newAllowedIPStore()
prefixes, ok := s.get(testPeer)
assert.False(t, ok, "an unconfigured peer must be reported as unknown, not as one without prefixes")
assert.Nil(t, prefixes, "an unknown peer has no prefixes")
}
func TestAllowedIPStoreAddUnions(t *testing.T) {
s := newAllowedIPStore()
overlay := netip.MustParsePrefix("100.64.0.1/32")
routed := netip.MustParsePrefix("10.20.0.0/16")
s.set(testPeer, []netip.Prefix{overlay})
// A peer update does not replace allowed IPs, and a repeated prefix must not be doubled.
s.add(testPeer, []netip.Prefix{overlay, routed})
prefixes, ok := s.get(testPeer)
require.True(t, ok, "peer must be known after set")
assert.Equal(t, []netip.Prefix{overlay, routed}, prefixes, "add must union rather than replace")
}
func TestAllowedIPStoreGetReturnsCopy(t *testing.T) {
s := newAllowedIPStore()
overlay := netip.MustParsePrefix("100.64.0.1/32")
s.set(testPeer, []netip.Prefix{overlay})
prefixes, ok := s.get(testPeer)
require.True(t, ok, "peer must be known after set")
prefixes[0] = netip.MustParsePrefix("0.0.0.0/0")
stored, _ := s.get(testPeer)
assert.Equal(t, []netip.Prefix{overlay}, stored, "a caller mutating the returned slice must not corrupt the store")
}
func TestAllowedIPStoreForgetAndReset(t *testing.T) {
s := newAllowedIPStore()
s.set(testPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32")})
s.set(otherPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.2/32")})
s.forget(testPeer)
_, ok := s.get(testPeer)
assert.False(t, ok, "a forgotten peer must be unknown")
_, ok = s.get(otherPeer)
assert.True(t, ok, "forgetting one peer must not touch the others")
s.reset()
_, ok = s.get(otherPeer)
assert.False(t, ok, "reset must drop every peer")
}
func TestIPNetsToPrefixes(t *testing.T) {
tests := []struct {
name string
ipNet net.IPNet
want string
}{
{
name: "v4",
ipNet: net.IPNet{IP: net.IP{10, 20, 0, 0}, Mask: net.CIDRMask(16, 32)},
want: "10.20.0.0/16",
},
{
name: "v4 mapped under a 128 bit mask",
ipNet: net.IPNet{IP: net.ParseIP("10.20.0.0"), Mask: net.CIDRMask(112, 128)},
want: "10.20.0.0/16",
},
{
name: "v6",
ipNet: net.IPNet{IP: net.ParseIP("fd00::"), Mask: net.CIDRMask(64, 128)},
want: "fd00::/64",
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
got := ipNetsToPrefixes([]net.IPNet{tc.ipNet})
require.Len(t, got, 1, "the address must be converted, not dropped")
assert.Equal(t, tc.want, got[0].String(), "converted prefix")
})
}
}
func TestIPNetsToPrefixesRoundTrip(t *testing.T) {
prefixes := []netip.Prefix{
netip.MustParsePrefix("100.64.0.1/32"),
netip.MustParsePrefix("10.20.0.0/16"),
netip.MustParsePrefix("fd00::/64"),
}
assert.Equal(t, prefixes, ipNetsToPrefixes(prefixesToIPNets(prefixes)),
"prefixes handed to a device must come back unchanged")
}
func TestAllowedIPStoreNormalizesMappedPrefixes(t *testing.T) {
s := newAllowedIPStore()
v4 := netip.MustParsePrefix("10.20.0.0/16")
mapped := netip.PrefixFrom(netip.AddrFrom16(v4.Addr().As16()), 112)
s.set(testPeer, []netip.Prefix{mapped})
// A v4 rule only matches a v4-mapped address once it has been unmapped, so the store must
// hold the plain form and recognise the two spellings as the same prefix.
s.add(testPeer, []netip.Prefix{v4})
prefixes, ok := s.get(testPeer)
require.True(t, ok, "peer must be known after set")
assert.Equal(t, []netip.Prefix{v4}, prefixes, "a mapped prefix must be stored unmapped and not duplicated")
}
func TestNormalizePrefix(t *testing.T) {
v4 := netip.MustParsePrefix("10.20.0.0/16")
v6 := netip.MustParsePrefix("fd00::/64")
assert.Equal(t, v4, normalizePrefix(v4), "a plain v4 prefix is unchanged")
assert.Equal(t, v6, normalizePrefix(v6), "a real v6 prefix is unchanged")
assert.Equal(t, v4, normalizePrefix(netip.PrefixFrom(netip.AddrFrom16(v4.Addr().As16()), 112)),
"a mapped prefix under a 128 bit mask becomes plain v4")
// A prefix shorter than /96 inside the mapped range is a genuine v6 prefix. Unmapping it
// would pair a v4 address with a v6 sized mask, which is invalid, and the store would then
// record a zero prefix that can never recreate the allowed IP.
for _, tc := range []string{"::ffff:0:0/64", "::ffff:1.2.3.4/80", "::ffff:1.2.3.4/95"} {
got := normalizePrefix(netip.MustParsePrefix(tc))
assert.True(t, got.IsValid(), "%s must normalize to a valid prefix", tc)
assert.False(t, got.Addr().Is4(), "%s must stay v6", tc)
}
}
func TestAllowedIPStoreAddExistingDoesNotCreate(t *testing.T) {
s := newAllowedIPStore()
routed := netip.MustParsePrefix("10.20.0.0/16")
// An update-only device operation on an absent peer is a silent no-op, so nothing may be
// recorded for a peer the store does not already know.
s.addExisting(testPeer, []netip.Prefix{routed})
_, ok := s.get(testPeer)
assert.False(t, ok, "addExisting must not record an unknown peer")
overlay := netip.MustParsePrefix("100.64.0.1/32")
s.set(testPeer, []netip.Prefix{overlay})
s.addExisting(testPeer, []netip.Prefix{routed})
prefixes, _ := s.get(testPeer)
assert.Equal(t, []netip.Prefix{overlay, routed}, prefixes, "addExisting must union onto a known peer")
}
func TestAllowedIPStoreHandsPrefixOverToTheNewOwner(t *testing.T) {
s := newAllowedIPStore()
routed := netip.MustParsePrefix("10.20.0.0/16")
other := otherPeer
s.set(testPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32"), routed})
s.set(other, []netip.Prefix{netip.MustParsePrefix("100.64.0.2/32")})
// The device takes an allowed IP away from its previous holder when it is configured on
// another peer, so the store must do the same rather than list it under both.
s.addExisting(other, []netip.Prefix{routed})
previous, _ := s.get(testPeer)
assert.NotContains(t, previous, routed, "the previous owner must lose the prefix")
current, _ := s.get(other)
assert.Contains(t, current, routed, "the new owner must hold the prefix")
}
func TestAllowedIPStoreForgetReleasesOwnership(t *testing.T) {
s := newAllowedIPStore()
routed := netip.MustParsePrefix("10.20.0.0/16")
s.set(testPeer, []netip.Prefix{routed})
s.forget(testPeer)
s.set(otherPeer, []netip.Prefix{routed})
// A forgotten peer must not be resurrected as a key in the peer map by a later claim.
_, ok := s.get(testPeer)
assert.False(t, ok, "the forgotten peer must stay unknown")
current, _ := s.get(otherPeer)
assert.Equal(t, []netip.Prefix{routed}, current, "the new owner must hold the prefix")
}
func TestNormalizePrefixClearsHostBits(t *testing.T) {
// A device stores a prefix masked, so a caller passing host bits must still match what a
// device fallback seeded, otherwise that prefix could never be removed by value.
assert.Equal(t, netip.MustParsePrefix("10.20.0.0/16"),
normalizePrefix(netip.MustParsePrefix("10.20.0.1/16")), "host bits must be cleared")
assert.Equal(t, netip.MustParsePrefix("fd00::/64"),
normalizePrefix(netip.MustParsePrefix("fd00::1/64")), "host bits must be cleared for v6")
}
func TestIPNetsToPrefixesKeepsV6InTheMappedRange(t *testing.T) {
// ::ffff:0:0/64 reads as v4-mapped but is a genuine v6 prefix: unmapping it would leave a
// v4 address under a 64 bit mask, which is invalid, and the allowed IP would be dropped.
got := ipNetsToPrefixes([]net.IPNet{{
IP: net.ParseIP("::ffff:0:0"),
Mask: net.CIDRMask(64, 128),
}})
require.Len(t, got, 1, "the prefix must be converted, not dropped")
assert.False(t, got[0].Addr().Is4(), "a v6 prefix in the mapped range must not become v4")
assert.Equal(t, 64, got[0].Bits(), "the prefix length must survive the conversion")
}
func TestPrefixesToIPNetsNormalizes(t *testing.T) {
// net.IPNet prints a v4-mapped address as v4 but takes the length from its 16 byte
// mask, so an unnormalized ::ffff:10.1.2.3/64 reaches a userspace device as 10.1.2.3/0,
// an allowed IP that matches every v4 address.
tests := []struct {
name string
given string
want string
}{
{name: "mapped below /96", given: "::ffff:10.1.2.3/64", want: "::/64"},
{name: "mapped at /112", given: "::ffff:10.1.2.3/112", want: "10.1.0.0/16"},
{name: "host bits are cleared", given: "10.20.0.1/16", want: "10.20.0.0/16"},
{name: "v6 is untouched", given: "fd00::1/64", want: "fd00::/64"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
got := prefixesToIPNets([]netip.Prefix{netip.MustParsePrefix(tc.given)})
require.Len(t, got, 1, "the prefix must be converted, not dropped")
assert.Equal(t, tc.want, got[0].String(), "what the device is given")
assert.NotEqual(t, 0, mustOnes(t, got[0]), "a device must never be given a zero length allowed IP")
})
}
}
func mustOnes(t *testing.T, ipNet net.IPNet) int {
t.Helper()
ones, _ := ipNet.Mask.Size()
return ones
}
// TestPrefixesToIPNetsAgreesWithTheStore pins the property the store depends on: what a
// device is given and what is recorded for it are the same prefix.
func TestPrefixesToIPNetsAgreesWithTheStore(t *testing.T) {
for _, given := range []string{"::ffff:10.1.2.3/64", "::ffff:10.1.2.3/112", "10.20.0.1/16", "fd00::1/64"} {
prefix := netip.MustParsePrefix(given)
toDevice := prefixesToIPNets([]netip.Prefix{prefix})
recorded := normalizePrefix(prefix)
assert.Equal(t, recorded.String(), toDevice[0].String(),
"%s must reach the device in the form the store records", given)
}
}
+8 -2
View File
@@ -19,12 +19,18 @@ func buildPresharedKeyConfig(peerKey wgtypes.Key, psk wgtypes.Key, updateOnly bo
} }
} }
// prefixesToIPNets converts prefixes on their way to a device. It is the only place that
// conversion happens, so it also normalizes: the device is then given the same form the
// store records, and a v4-mapped prefix cannot reach net.IPNet, which prints such an
// address as v4 while taking the length from its 16 byte mask and so turns
// ::ffff:10.1.2.3/64 into 10.1.2.3/0 — an allowed IP matching every v4 address.
func prefixesToIPNets(prefixes []netip.Prefix) []net.IPNet { func prefixesToIPNets(prefixes []netip.Prefix) []net.IPNet {
ipNets := make([]net.IPNet, len(prefixes)) ipNets := make([]net.IPNet, len(prefixes))
for i, prefix := range prefixes { for i, prefix := range prefixes {
normalized := normalizePrefix(prefix)
ipNets[i] = net.IPNet{ ipNets[i] = net.IPNet{
IP: prefix.Addr().AsSlice(), // Convert netip.Addr to net.IP IP: normalized.Addr().AsSlice(),
Mask: net.CIDRMask(prefix.Bits(), prefix.Addr().BitLen()), // Create subnet mask Mask: net.CIDRMask(normalized.Bits(), normalized.Addr().BitLen()),
} }
} }
return ipNets return ipNets
+74 -34
View File
@@ -6,6 +6,7 @@ import (
"fmt" "fmt"
"net" "net"
"net/netip" "net/netip"
"slices"
"time" "time"
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
@@ -18,16 +19,22 @@ import (
type KernelConfigurer struct { type KernelConfigurer struct {
deviceName string deviceName string
statsCache *statsCache statsCache *statsCache
allowedIPs *allowedIPStore
} }
// NewKernelConfigurer creates a configurer with an empty allowed IP mirror
// and a statistics cache for the named kernel device.
func NewKernelConfigurer(deviceName string) *KernelConfigurer { func NewKernelConfigurer(deviceName string) *KernelConfigurer {
c := &KernelConfigurer{ c := &KernelConfigurer{
deviceName: deviceName, deviceName: deviceName,
allowedIPs: newAllowedIPStore(),
} }
c.statsCache = newStatsCache(statsCacheTTL, c.fetchStats) c.statsCache = newStatsCache(statsCacheTTL, c.fetchStats)
return c return c
} }
// ConfigureInterface sets the device key, port and firewall mark, replacing all peers.
// The allowed IP mirror is reset only after the device accepts the configuration.
func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error { func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error {
log.Debugf("adding Wireguard private key") log.Debugf("adding Wireguard private key")
key, err := wgtypes.ParseKey(privateKey) key, err := wgtypes.ParseKey(privateKey)
@@ -46,6 +53,8 @@ func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error
if err != nil { if err != nil {
return fmt.Errorf(`received error "%w" while configuring interface %s with port %d`, err, c.deviceName, port) return fmt.Errorf(`received error "%w" while configuring interface %s with port %d`, err, c.deviceName, port)
} }
c.allowedIPs.reset()
return nil return nil
} }
@@ -58,9 +67,20 @@ func (c *KernelConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, upda
} }
cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly) cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly)
return c.configure(cfg) if err := c.configure(cfg); err != nil {
return err
}
// Without updateOnly this creates the peer when it is absent, so the store has to
// know about it even though no allowed IP was configured.
if !updateOnly {
c.allowedIPs.ensure(parsedPeerKey)
}
return nil
} }
// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set.
// Prefixes assigned to this peer are transferred from their previous owners.
func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error { func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey) peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil { if err != nil {
@@ -83,19 +103,23 @@ func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix,
if err != nil { if err != nil {
return fmt.Errorf(`received error "%w" while updating peer on interface %s with settings: allowed ips %s, endpoint %s`, err, c.deviceName, allowedIps, endpoint.String()) return fmt.Errorf(`received error "%w" while updating peer on interface %s with settings: allowed ips %s, endpoint %s`, err, c.deviceName, allowedIps, endpoint.String())
} }
c.allowedIPs.add(peerKeyParsed, allowedIps)
return nil return nil
} }
// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured.
// Neither the netlink API nor the userspace one can clear an endpoint in place, so the peer
// is removed and re-added with the allowed IPs it already had.
func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error { func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey) peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil { if err != nil {
return err return err
} }
// Get the existing peer to preserve its allowed IPs allowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
existingPeer, err := c.getPeer(c.deviceName, peerKey)
if err != nil { if err != nil {
return fmt.Errorf("get peer: %w", err) return err
} }
removePeerCfg := wgtypes.PeerConfig{ removePeerCfg := wgtypes.PeerConfig{
@@ -104,26 +128,27 @@ func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error {
} }
if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{removePeerCfg}}); err != nil { if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{removePeerCfg}}); err != nil {
return fmt.Errorf(`error removing peer %s from interface %s: %w`, peerKey, c.deviceName, err) return fmt.Errorf("remove peer %s from interface %s: %w", peerKey, c.deviceName, err)
} }
//Re-add the peer without the endpoint but same AllowedIPs
reAddPeerCfg := wgtypes.PeerConfig{ reAddPeerCfg := wgtypes.PeerConfig{
PublicKey: peerKeyParsed, PublicKey: peerKeyParsed,
AllowedIPs: existingPeer.AllowedIPs, AllowedIPs: prefixesToIPNets(allowedIPs),
ReplaceAllowedIPs: true, ReplaceAllowedIPs: true,
} }
if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{reAddPeerCfg}}); err != nil { if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{reAddPeerCfg}}); err != nil {
c.allowedIPs.forget(peerKeyParsed)
return fmt.Errorf( return fmt.Errorf(
`error re-adding peer %s to interface %s with allowed IPs %v: %w`, "re-add peer %s to interface %s with allowed IPs %v: %w",
peerKey, c.deviceName, existingPeer.AllowedIPs, err, peerKey, c.deviceName, allowedIPs, err,
) )
} }
return nil return nil
} }
// RemovePeer removes a peer and forgets its allowed IPs after a successful device write.
func (c *KernelConfigurer) RemovePeer(peerKey string) error { func (c *KernelConfigurer) RemovePeer(peerKey string) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey) peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil { if err != nil {
@@ -142,15 +167,13 @@ func (c *KernelConfigurer) RemovePeer(peerKey string) error {
if err != nil { if err != nil {
return fmt.Errorf(`received error "%w" while removing peer %s from interface %s`, err, peerKey, c.deviceName) return fmt.Errorf(`received error "%w" while removing peer %s from interface %s`, err, peerKey, c.deviceName)
} }
c.allowedIPs.forget(peerKeyParsed)
return nil return nil
} }
// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op.
func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error { func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
ipNet := net.IPNet{
IP: allowedIP.Addr().AsSlice(),
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
}
peerKeyParsed, err := wgtypes.ParseKey(peerKey) peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil { if err != nil {
return err return err
@@ -159,7 +182,7 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix)
PublicKey: peerKeyParsed, PublicKey: peerKeyParsed,
UpdateOnly: true, UpdateOnly: true,
ReplaceAllowedIPs: false, ReplaceAllowedIPs: false,
AllowedIPs: []net.IPNet{ipNet}, AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}),
} }
config := wgtypes.Config{ config := wgtypes.Config{
@@ -169,52 +192,69 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix)
if err != nil { if err != nil {
return fmt.Errorf(`received error "%w" while adding allowed Ip to peer on interface %s with settings: allowed ips %s`, err, c.deviceName, allowedIP) return fmt.Errorf(`received error "%w" while adding allowed Ip to peer on interface %s with settings: allowed ips %s`, err, c.deviceName, allowedIP)
} }
c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP})
return nil return nil
} }
// RemoveAllowedIP removes a prefix while preserving the peer's other allowed IPs.
// A prefix not assigned to the peer is a no-op.
func (c *KernelConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error { func (c *KernelConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error {
ipNet := net.IPNet{
IP: allowedIP.Addr().AsSlice(),
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
}
peerKeyParsed, err := wgtypes.ParseKey(peerKey) peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil { if err != nil {
return fmt.Errorf("parse peer key: %w", err) return fmt.Errorf("parse peer key: %w", err)
} }
existingPeer, err := c.getPeer(c.deviceName, peerKey) currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
if err != nil { if err != nil {
return fmt.Errorf("get peer: %w", err) return err
} }
newAllowedIPs := existingPeer.AllowedIPs idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP))
if idx < 0 {
for i, existingAllowedIP := range existingPeer.AllowedIPs { return nil
if existingAllowedIP.String() == ipNet.String() {
newAllowedIPs = append(existingPeer.AllowedIPs[:i], existingPeer.AllowedIPs[i+1:]...) //nolint:gocritic
break
}
} }
newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1)
peer := wgtypes.PeerConfig{ peer := wgtypes.PeerConfig{
PublicKey: peerKeyParsed, PublicKey: peerKeyParsed,
UpdateOnly: true, UpdateOnly: true,
ReplaceAllowedIPs: true, ReplaceAllowedIPs: true,
AllowedIPs: newAllowedIPs, AllowedIPs: prefixesToIPNets(newAllowedIPs),
} }
config := wgtypes.Config{ config := wgtypes.Config{
Peers: []wgtypes.PeerConfig{peer}, Peers: []wgtypes.PeerConfig{peer},
} }
err = c.configure(config) if err := c.configure(config); err != nil {
if err != nil {
return fmt.Errorf("remove allowed IP %s on interface %s: %w", allowedIP, c.deviceName, err) return fmt.Errorf("remove allowed IP %s on interface %s: %w", allowedIP, c.deviceName, err)
} }
c.allowedIPs.set(peerKeyParsed, newAllowedIPs)
return nil return nil
} }
func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) { // peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device
// only for a peer the store has not seen. Dumping the device costs a netlink round trip
// proportional to the whole network map, and this runs on every relay and ICE transition.
func (c *KernelConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) {
if prefixes, ok := c.allowedIPs.get(peerKey); ok {
return prefixes, nil
}
existingPeer, err := c.getPeer(c.deviceName, peerKey)
if err != nil {
return nil, fmt.Errorf("get peer: %w", err)
}
prefixes := ipNetsToPrefixes(existingPeer.AllowedIPs)
c.allowedIPs.set(peerKey, prefixes)
return prefixes, nil
}
// getPeer scans the device for one peer. wgtypes.Key is an array, so the comparison is a
// plain equality: Key.String would base64 encode into a fresh allocation for every peer.
func (c *KernelConfigurer) getPeer(ifaceName string, peerPubKey wgtypes.Key) (wgtypes.Peer, error) {
wg, err := wgctrl.New() wg, err := wgctrl.New()
if err != nil { if err != nil {
return wgtypes.Peer{}, fmt.Errorf("wgctl: %w", err) return wgtypes.Peer{}, fmt.Errorf("wgctl: %w", err)
@@ -231,7 +271,7 @@ func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer,
return wgtypes.Peer{}, fmt.Errorf("get device %s: %w", ifaceName, err) return wgtypes.Peer{}, fmt.Errorf("get device %s: %w", ifaceName, err)
} }
for _, peer := range wgDevice.Peers { for _, peer := range wgDevice.Peers {
if peer.PublicKey.String() == peerPubKey { if peer.PublicKey == peerPubKey {
return peer, nil return peer, nil
} }
} }
+120 -92
View File
@@ -8,6 +8,7 @@ import (
"net/netip" "net/netip"
"os" "os"
"runtime" "runtime"
"slices"
"strconv" "strconv"
"strings" "strings"
"time" "time"
@@ -41,31 +42,38 @@ type WGUSPConfigurer struct {
deviceName string deviceName string
activityRecorder *bind.ActivityRecorder activityRecorder *bind.ActivityRecorder
statsCache *statsCache statsCache *statsCache
allowedIPs *allowedIPStore
uapiListener net.Listener uapiListener net.Listener
} }
// NewUSPConfigurer creates a userspace configurer and starts its UAPI listener.
func NewUSPConfigurer(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer { func NewUSPConfigurer(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
wgCfg := &WGUSPConfigurer{ wgCfg := &WGUSPConfigurer{
device: device, device: device,
deviceName: deviceName, deviceName: deviceName,
activityRecorder: activityRecorder, activityRecorder: activityRecorder,
allowedIPs: newAllowedIPStore(),
} }
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats) wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
wgCfg.startUAPI() wgCfg.startUAPI()
return wgCfg return wgCfg
} }
// NewUSPConfigurerNoUAPI creates a userspace configurer without a UAPI listener.
func NewUSPConfigurerNoUAPI(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer { func NewUSPConfigurerNoUAPI(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
wgCfg := &WGUSPConfigurer{ wgCfg := &WGUSPConfigurer{
device: device, device: device,
deviceName: deviceName, deviceName: deviceName,
activityRecorder: activityRecorder, activityRecorder: activityRecorder,
allowedIPs: newAllowedIPStore(),
} }
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats) wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
return wgCfg return wgCfg
} }
// ConfigureInterface sets the device key, port and firewall mark, replacing all peers.
// The allowed IP mirror is reset only after the device accepts the configuration.
func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error { func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error {
log.Debugf("adding Wireguard private key") log.Debugf("adding Wireguard private key")
key, err := wgtypes.ParseKey(privateKey) key, err := wgtypes.ParseKey(privateKey)
@@ -80,7 +88,12 @@ func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error
ListenPort: &port, ListenPort: &port,
} }
return c.device.IpcSet(toWgUserspaceString(config)) if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
return err
}
c.allowedIPs.reset()
return nil
} }
// SetPresharedKey sets the preshared key for a peer. // SetPresharedKey sets the preshared key for a peer.
@@ -92,14 +105,38 @@ func (c *WGUSPConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, updat
} }
cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly) cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly)
return c.device.IpcSet(toWgUserspaceString(cfg)) if err := c.device.IpcSet(toWgUserspaceString(cfg)); err != nil {
return err
}
// Without updateOnly this creates the peer when it is absent, so the store has to
// know about it even though no allowed IP was configured.
if !updateOnly {
c.allowedIPs.ensure(parsedPeerKey)
}
return nil
} }
// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set.
// It validates the endpoint before writing and records changes after a successful write.
func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error { func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey) peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil { if err != nil {
return err return err
} }
// Everything that can fail is done before the device is touched, so a failure here
// cannot leave the device holding a peer that the activity recorder and the allowed
// IP store never learned about.
var addrPort netip.AddrPort
if endpoint != nil {
addr, err := netip.ParseAddr(endpoint.IP.String())
if err != nil {
return fmt.Errorf("parse endpoint address: %w", err)
}
addrPort = netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port))
}
peer := wgtypes.PeerConfig{ peer := wgtypes.PeerConfig{
PublicKey: peerKeyParsed, PublicKey: peerKeyParsed,
ReplaceAllowedIPs: false, ReplaceAllowedIPs: false,
@@ -119,47 +156,27 @@ func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix,
} }
if endpoint != nil { if endpoint != nil {
addr, err := netip.ParseAddr(endpoint.IP.String())
if err != nil {
return fmt.Errorf("failed to parse endpoint address: %w", err)
}
addrPort := netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port))
c.activityRecorder.UpsertAddress(peerKey, addrPort) c.activityRecorder.UpsertAddress(peerKey, addrPort)
} }
c.allowedIPs.add(peerKeyParsed, allowedIps)
return nil return nil
} }
// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured.
// The UAPI cannot clear an endpoint in place, so the peer is removed and re-added with the
// allowed IPs it already had.
func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error { func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey) peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil { if err != nil {
return fmt.Errorf("parse peer key: %w", err) return fmt.Errorf("parse peer key: %w", err)
} }
ipcStr, err := c.device.IpcGet() allowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
if err != nil { if err != nil {
return fmt.Errorf("get IPC config: %w", err) return err
} }
// Parse current status to get allowed IPs for the peer
stats, err := parseStatus(c.deviceName, ipcStr)
if err != nil {
return fmt.Errorf("parse IPC config: %w", err)
}
var allowedIPs []net.IPNet
found := false
for _, peer := range stats.Peers {
if peer.PublicKey == peerKey {
allowedIPs = peer.AllowedIPs
found = true
break
}
}
if !found {
return fmt.Errorf("peer %s not found", peerKey)
}
// remove the peer from the WireGuard configuration
peer := wgtypes.PeerConfig{ peer := wgtypes.PeerConfig{
PublicKey: peerKeyParsed, PublicKey: peerKeyParsed,
Remove: true, Remove: true,
@@ -169,14 +186,13 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
Peers: []wgtypes.PeerConfig{peer}, Peers: []wgtypes.PeerConfig{peer},
} }
if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil { if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil {
return fmt.Errorf("failed to remove peer: %s", ipcErr) return fmt.Errorf("remove peer: %w", ipcErr)
} }
// Build the peer config
peer = wgtypes.PeerConfig{ peer = wgtypes.PeerConfig{
PublicKey: peerKeyParsed, PublicKey: peerKeyParsed,
ReplaceAllowedIPs: true, ReplaceAllowedIPs: true,
AllowedIPs: allowedIPs, AllowedIPs: prefixesToIPNets(allowedIPs),
} }
config = wgtypes.Config{ config = wgtypes.Config{
@@ -184,12 +200,15 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
} }
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil { if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
return fmt.Errorf("remove endpoint address: %w", err) c.allowedIPs.forget(peerKeyParsed)
return fmt.Errorf("re-add peer without endpoint: %w", err)
} }
return nil return nil
} }
// RemovePeer removes a peer, then clears its activity and allowed IP records.
// A failed device write leaves both records intact.
func (c *WGUSPConfigurer) RemovePeer(peerKey string) error { func (c *WGUSPConfigurer) RemovePeer(peerKey string) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey) peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil { if err != nil {
@@ -204,18 +223,17 @@ func (c *WGUSPConfigurer) RemovePeer(peerKey string) error {
config := wgtypes.Config{ config := wgtypes.Config{
Peers: []wgtypes.PeerConfig{peer}, Peers: []wgtypes.PeerConfig{peer},
} }
ipcErr := c.device.IpcSet(toWgUserspaceString(config)) if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil {
return ipcErr
c.activityRecorder.Remove(peerKey)
return ipcErr
}
func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
ipNet := net.IPNet{
IP: allowedIP.Addr().AsSlice(),
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
} }
c.activityRecorder.Remove(peerKey)
c.allowedIPs.forget(peerKeyParsed)
return nil
}
// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op.
func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey) peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil { if err != nil {
return err return err
@@ -224,79 +242,89 @@ func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) e
PublicKey: peerKeyParsed, PublicKey: peerKeyParsed,
UpdateOnly: true, UpdateOnly: true,
ReplaceAllowedIPs: false, ReplaceAllowedIPs: false,
AllowedIPs: []net.IPNet{ipNet}, AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}),
} }
config := wgtypes.Config{ config := wgtypes.Config{
Peers: []wgtypes.PeerConfig{peer}, Peers: []wgtypes.PeerConfig{peer},
} }
return c.device.IpcSet(toWgUserspaceString(config)) if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
return err
}
c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP})
return nil
} }
// RemoveAllowedIP removes a prefix while preserving the peer's other allowed IPs.
// It returns ErrAllowedIPNotFound if the prefix is not assigned to the peer.
func (c *WGUSPConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error { func (c *WGUSPConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error {
ipc, err := c.device.IpcGet()
if err != nil {
return err
}
peerKeyParsed, err := wgtypes.ParseKey(peerKey) peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return fmt.Errorf("parse peer key: %w", err)
}
currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
if err != nil { if err != nil {
return err return err
} }
hexKey := hex.EncodeToString(peerKeyParsed[:])
lines := strings.Split(ipc, "\n") idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP))
if idx < 0 {
return ErrAllowedIPNotFound
}
newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1)
peer := wgtypes.PeerConfig{ peer := wgtypes.PeerConfig{
PublicKey: peerKeyParsed, PublicKey: peerKeyParsed,
UpdateOnly: true, UpdateOnly: true,
ReplaceAllowedIPs: true, ReplaceAllowedIPs: true,
AllowedIPs: []net.IPNet{}, AllowedIPs: prefixesToIPNets(newAllowedIPs),
} }
foundPeer := false
removedAllowedIP := false
ip := allowedIP.String()
for _, line := range lines {
line = strings.TrimSpace(line)
// If we're within the details of the found peer and encounter another public key,
// this means we're starting another peer's details. So, reset the flag.
if strings.HasPrefix(line, "public_key=") && foundPeer {
foundPeer = false
}
// Identify the peer with the specific public key
if line == fmt.Sprintf("public_key=%s", hexKey) {
foundPeer = true
}
// If we're within the details of the found peer and find the specific allowed IP, skip this line
if foundPeer && line == "allowed_ip="+ip {
removedAllowedIP = true
continue
}
// Append the line to the output string
if foundPeer && strings.HasPrefix(line, "allowed_ip=") {
allowedIPStr := strings.TrimPrefix(line, "allowed_ip=")
_, ipNet, err := net.ParseCIDR(allowedIPStr)
if err != nil {
return err
}
peer.AllowedIPs = append(peer.AllowedIPs, *ipNet)
}
}
if !removedAllowedIP {
return ErrAllowedIPNotFound
}
config := wgtypes.Config{ config := wgtypes.Config{
Peers: []wgtypes.PeerConfig{peer}, Peers: []wgtypes.PeerConfig{peer},
} }
return c.device.IpcSet(toWgUserspaceString(config)) if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
return fmt.Errorf("remove allowed IP %s: %w", allowedIP, err)
}
c.allowedIPs.set(peerKeyParsed, newAllowedIPs)
return nil
}
// peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device
// only for a peer the store has not seen. Reading them back means dumping and parsing the
// whole device configuration, and this runs on every relay and ICE transition.
func (c *WGUSPConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) {
if prefixes, ok := c.allowedIPs.get(peerKey); ok {
return prefixes, nil
}
ipcStr, err := c.device.IpcGet()
if err != nil {
return nil, fmt.Errorf("get IPC config: %w", err)
}
stats, err := parseStatus(c.deviceName, ipcStr)
if err != nil {
return nil, fmt.Errorf("parse IPC config: %w", err)
}
// parseStatus reports keys in their textual form, so the comparison needs it once.
wanted := peerKey.String()
for _, peer := range stats.Peers {
if peer.PublicKey != wanted {
continue
}
prefixes := ipNetsToPrefixes(peer.AllowedIPs)
c.allowedIPs.set(peerKey, prefixes)
return prefixes, nil
}
return nil, ErrPeerNotFound
} }
func (c *WGUSPConfigurer) FullStats() (*Stats, error) { func (c *WGUSPConfigurer) FullStats() (*Stats, error) {
@@ -0,0 +1,318 @@
package configurer
import (
"net"
"net/netip"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
wgconn "golang.zx2c4.com/wireguard/conn"
wgdevice "golang.zx2c4.com/wireguard/device"
"golang.zx2c4.com/wireguard/tun/tuntest"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"github.com/netbirdio/netbird/client/iface/bind"
)
// newTestUSPConfigurer builds a configurer over a real wireguard-go device backed by an
// in-memory TUN. The device stays down, so no socket is opened and no privileges are needed.
func newTestUSPConfigurer(t *testing.T) *WGUSPConfigurer {
t.Helper()
tun := tuntest.NewChannelTUN()
dev := wgdevice.NewDevice(tun.TUN(), wgconn.NewDefaultBind(), wgdevice.NewLogger(wgdevice.LogLevelSilent, ""))
t.Cleanup(dev.Close)
c := NewUSPConfigurerNoUAPI(dev, "wgtest0", bind.NewActivityRecorder())
key, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate device private key")
require.NoError(t, c.ConfigureInterface(key.String(), 0), "configure test device")
return c
}
// seedPeers adds count peers, each with a /32 overlay address, and returns their public keys.
func seedPeers(t *testing.T, c *WGUSPConfigurer, count int) []string {
t.Helper()
keys := make([]string, 0, count)
for i := 0; i < count; i++ {
priv, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate peer private key")
pub := priv.PublicKey().String()
addr := netip.PrefixFrom(netip.AddrFrom4([4]byte{100, 64, byte(i >> 8), byte(i)}), 32)
require.NoError(t, c.UpdatePeer(pub, []netip.Prefix{addr}, 25*time.Second, nil, nil), "add peer")
keys = append(keys, pub)
}
return keys
}
func peerAllowedIPs(t *testing.T, c *WGUSPConfigurer, peerKey string) []string {
t.Helper()
stats, err := c.FullStats()
require.NoError(t, err, "read device stats")
for _, p := range stats.Peers {
if p.PublicKey != peerKey {
continue
}
got := make([]string, 0, len(p.AllowedIPs))
for _, ipNet := range p.AllowedIPs {
got = append(got, ipNet.String())
}
return got
}
t.Fatalf("peer %s not found on device", peerKey)
return nil
}
// TestRemoveEndpointAddressPreservesRoutedAllowedIPs covers the prefixes the route manager
// attaches to a routing peer through AddAllowedIP. Those are not known to the peer.Conn that
// triggers the endpoint removal, so dropping them here would silently blackhole every route
// behind that peer on each relay or ICE disconnect.
func TestRemoveEndpointAddressPreservesRoutedAllowedIPs(t *testing.T) {
c := newTestUSPConfigurer(t)
peerKey := seedPeers(t, c, 3)[1]
routed := []netip.Prefix{
netip.MustParsePrefix("10.20.0.0/16"),
netip.MustParsePrefix("192.168.7.0/24"),
}
for _, prefix := range routed {
require.NoError(t, c.AddAllowedIP(peerKey, prefix), "add routed prefix")
}
before := peerAllowedIPs(t, c, peerKey)
require.Len(t, before, 3, "peer should hold its overlay address plus both routed prefixes")
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey),
"allowed IPs must survive the endpoint removal unchanged")
}
// TestRemoveEndpointAddressDoesNotScaleWithPeerCount is the regression guard for the actual
// defect: clearing one peer's endpoint used to dump and parse the whole device, so its cost
// grew with the size of the network map. On a routing peer with thousands of peers that dump
// runs on every relay and ICE transition, under the interface lock.
func TestRemoveEndpointAddressDoesNotScaleWithPeerCount(t *testing.T) {
measure := func(peerCount int) float64 {
c := newTestUSPConfigurer(t)
peerKey := seedPeers(t, c, peerCount)[peerCount/2]
return testing.AllocsPerRun(5, func() {
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
})
}
small := measure(64)
large := measure(1024)
assert.Less(t, large, small*2,
"clearing one endpoint allocated %.0f objects with 1024 peers against %.0f with 64: the cost still scales with the peer count",
large, small)
}
// TestRemoveEndpointAddressFallsBackToDevice covers a peer the store never saw, which is what
// an out-of-band reconfiguration of the device leaves behind. The device stays the source of
// truth in that case, so the allowed IPs must still be preserved.
func TestRemoveEndpointAddressFallsBackToDevice(t *testing.T) {
c := newTestUSPConfigurer(t)
peerKey := seedPeers(t, c, 3)[1]
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix")
before := peerAllowedIPs(t, c, peerKey)
c.allowedIPs.reset()
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey),
"allowed IPs recovered from the device must be preserved")
recovered, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
assert.True(t, ok, "the fallback must seed the store so the next call skips the device dump")
assert.Len(t, recovered, 2, "seeded prefixes")
}
func TestRemoveAllowedIPKeepsTheOtherPrefixes(t *testing.T) {
c := newTestUSPConfigurer(t)
peerKey := seedPeers(t, c, 3)[0]
routed := netip.MustParsePrefix("10.20.0.0/16")
require.NoError(t, c.AddAllowedIP(peerKey, routed), "add routed prefix")
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("192.168.7.0/24")), "add routed prefix")
require.NoError(t, c.RemoveAllowedIP(peerKey, routed), "remove routed prefix")
assert.ElementsMatch(t, []string{"100.64.0.0/32", "192.168.7.0/24"}, peerAllowedIPs(t, c, peerKey),
"only the removed prefix should be gone")
assert.ErrorIs(t, c.RemoveAllowedIP(peerKey, routed), ErrAllowedIPNotFound,
"removing a prefix that is no longer configured must be reported")
}
// TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt covers the lazy connection window documented
// in #6863: AddAllowedIP is update-only, a silent no-op when the peer is absent, so it must not
// leave the store claiming prefixes the device never took. RemoveEndpointAddress re-adds a peer
// without update-only, so a phantom entry would create a peer the device had dropped, and a
// created peer would steal those allowed IPs from whichever peer legitimately holds them.
func TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt(t *testing.T) {
c := newTestUSPConfigurer(t)
seedPeers(t, c, 2)
priv, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate peer private key")
absent := priv.PublicKey().String()
require.NoError(t, c.AddAllowedIP(absent, netip.MustParsePrefix("10.20.0.0/16")),
"update-only add on an absent peer is a silent no-op")
stats, err := c.FullStats()
require.NoError(t, err, "read device stats")
require.Len(t, stats.Peers, 2, "the absent peer must not have been created by AddAllowedIP")
assert.ErrorIs(t, c.RemoveEndpointAddress(absent), ErrPeerNotFound,
"clearing the endpoint of a peer the device does not have must fail")
stats, err = c.FullStats()
require.NoError(t, err, "read device stats")
assert.Len(t, stats.Peers, 2, "no peer may be created while clearing an endpoint")
}
// TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer covers WireGuard's rule that an
// allowed IP belongs to exactly one peer: configuring a prefix on a peer takes it away from
// whichever peer held it before. UpdatePeer relies on that rule rather than removing the prefix
// from the previous holder itself, so a prefix handed over between peers must not come back.
func TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer(t *testing.T) {
c := newTestUSPConfigurer(t)
keys := seedPeers(t, c, 2)
peerA, peerB := keys[0], keys[1]
routed := netip.MustParsePrefix("10.20.0.0/16")
require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A")
require.Contains(t, peerAllowedIPs(t, c, peerA), routed.String(), "A must hold the prefix")
// The route moves to B. The device takes it away from A on its own.
require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B")
require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix")
require.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(), "the device must have taken it from A")
require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint")
assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(),
"clearing A's endpoint must not take the prefix back from B")
assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(),
"B must still hold the prefix")
}
// TestPresharedKeyCreatedPeerTakesPartInPrefixHandover covers a peer created by a preshared
// key write rather than by a peer update. Rosenpass applies a peer's first key without
// updateOnly, which creates the peer on the device, so a store that ignored that operation
// would treat the peer as unknown and would not account for a prefix later handed over to it.
func TestPresharedKeyCreatedPeerTakesPartInPrefixHandover(t *testing.T) {
c := newTestUSPConfigurer(t)
peerA := seedPeers(t, c, 1)[0]
routed := netip.MustParsePrefix("10.20.0.0/16")
require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A")
priv, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate peer private key")
peerB := priv.PublicKey().String()
psk, err := wgtypes.GenerateKey()
require.NoError(t, err, "generate preshared key")
require.NoError(t, c.SetPresharedKey(peerB, psk, false), "a first key creates the peer")
require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B")
require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix")
require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint")
assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(),
"clearing A's endpoint must not take the prefix back from B")
assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must still hold the prefix")
}
// TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice is the end to end form of the
// conversion: a v4-mapped prefix must not reach the device as a zero length allowed IP,
// which would route every v4 address to that peer.
func TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice(t *testing.T) {
c := newTestUSPConfigurer(t)
priv, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate peer private key")
peerKey := priv.PublicKey().String()
mapped := netip.MustParsePrefix("::ffff:10.1.2.3/112")
require.NoError(t, c.UpdatePeer(peerKey, []netip.Prefix{mapped}, 25*time.Second, nil, nil), "add peer")
onDevice := peerAllowedIPs(t, c, peerKey)
assert.NotContains(t, onDevice, "0.0.0.0/0", "the device must not be given a catch-all allowed IP")
assert.Equal(t, []string{"10.1.0.0/16"}, onDevice, "the device holds the normalized prefix")
recorded, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
require.True(t, ok, "the peer must be recorded")
require.Len(t, recorded, 1, "one prefix recorded")
assert.Equal(t, onDevice[0], recorded[0].String(), "device and store must agree")
}
// TestUpdatePeerWithAnUnusableEndpointTouchesNothing pins the ordering: the endpoint is
// parsed before the device is configured, so a failure cannot leave the device holding a
// peer that the store never learned about, with the prefix handover skipped along with it.
func TestUpdatePeerWithAnUnusableEndpointTouchesNothing(t *testing.T) {
c := newTestUSPConfigurer(t)
seedPeers(t, c, 2)
priv, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate peer private key")
peerKey := priv.PublicKey().String()
// A three byte address has no textual form netip can parse back.
endpoint := &net.UDPAddr{IP: net.IP{1, 2, 3}, Port: 51820}
require.Error(t, c.UpdatePeer(peerKey, []netip.Prefix{netip.MustParsePrefix("10.30.0.0/16")},
25*time.Second, endpoint, nil), "an unusable endpoint must fail the update")
stats, err := c.FullStats()
require.NoError(t, err, "read device stats")
assert.Len(t, stats.Peers, 2, "the peer must not have reached the device")
_, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
assert.False(t, ok, "the peer must not have been recorded either")
}
// TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses covers a removal that never reached the
// device. A single peer removal is one write, so a failure leaves the peer on the device
// exactly as it was, and the record still describes it; dropping it would only force the
// next caller to read the whole device back for an answer it already had.
func TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses(t *testing.T) {
c := newTestUSPConfigurer(t)
peerKey := seedPeers(t, c, 1)[0]
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix")
before, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
require.True(t, ok, "the peer must be recorded before the removal")
require.Len(t, before, 2, "overlay address plus routed prefix")
// A closed device refuses every write, which is the shape of any failed removal.
c.device.Close()
require.Error(t, c.RemovePeer(peerKey), "the removal must report the failure")
after, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
require.True(t, ok, "a peer still on the device must stay recorded")
assert.Equal(t, before, after, "the record must describe the peer the device kept")
}
// mustParseKey turns the textual key the configurer API takes into the form the store
// keys on.
func mustParseKey(t *testing.T, key string) wgtypes.Key {
t.Helper()
parsed, err := wgtypes.ParseKey(key)
require.NoError(t, err, "parse peer key")
return parsed
}
-7
View File
@@ -51,7 +51,6 @@ func ValidateMTU(mtu uint16) error {
type wgProxyFactory interface { type wgProxyFactory interface {
GetProxy() wgproxy.Proxy GetProxy() wgproxy.Proxy
GetProxyPort() uint16
Free() error Free() error
} }
@@ -81,12 +80,6 @@ func (w *WGIface) GetProxy() wgproxy.Proxy {
return w.wgProxyFactory.GetProxy() return w.wgProxyFactory.GetProxy()
} }
// GetProxyPort returns the proxy port used by the WireGuard proxy.
// Returns 0 if no proxy port is used (e.g., for userspace WireGuard).
func (w *WGIface) GetProxyPort() uint16 {
return w.wgProxyFactory.GetProxyPort()
}
// GetBind returns the EndpointManager userspace bind mode. // GetBind returns the EndpointManager userspace bind mode.
func (w *WGIface) GetBind() device.EndpointManager { func (w *WGIface) GetBind() device.EndpointManager {
w.mu.Lock() w.mu.Lock()
-1
View File
@@ -52,7 +52,6 @@ func (f *fakeTunDevice) Close() error {
type fakeProxyFactory struct{} type fakeProxyFactory struct{}
func (fakeProxyFactory) GetProxy() wgproxy.Proxy { return nil } func (fakeProxyFactory) GetProxy() wgproxy.Proxy { return nil }
func (fakeProxyFactory) GetProxyPort() uint16 { return 0 }
func (fakeProxyFactory) Free() error { return nil } func (fakeProxyFactory) Free() error { return nil }
// TestWGIface_CloseReleasesMutexBeforeTunClose guards against a deadlock // TestWGIface_CloseReleasesMutexBeforeTunClose guards against a deadlock
+2 -15
View File
@@ -6,27 +6,14 @@ import (
"fmt" "fmt"
"os/exec" "os/exec"
log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/client/internal/wincmd"
) )
func (w *WGIface) Destroy() error { func (w *WGIface) Destroy() error {
netshCmd := GetSystem32Command("netsh") netshCmd := wincmd.System32("netsh")
out, err := exec.Command(netshCmd, "interface", "set", "interface", w.Name(), "admin=disable").CombinedOutput() out, err := exec.Command(netshCmd, "interface", "set", "interface", w.Name(), "admin=disable").CombinedOutput()
if err != nil { if err != nil {
return fmt.Errorf("failed to remove interface %s: %w - %s", w.Name(), err, out) return fmt.Errorf("failed to remove interface %s: %w - %s", w.Name(), err, out)
} }
return nil return nil
} }
// GetSystem32Command checks if a command can be found in the system path and returns it. In case it can't find it
// in the path it will return the full path of a command assuming C:\windows\system32 as the base path.
func GetSystem32Command(command string) string {
_, err := exec.LookPath(command)
if err == nil {
return command
}
log.Tracef("Command %s not found in PATH, using C:\\windows\\system32\\%s.exe path", command, command)
return "C:\\windows\\system32\\" + command + ".exe"
}
+38 -42
View File
@@ -40,14 +40,18 @@ func init() {
peerPubKey = peerPrivateKey.PublicKey().String() peerPubKey = peerPrivateKey.PublicKey().String()
} }
// testIFaceBlackList mirrors the prefixes profilemanager.DefaultInterfaceBlacklist
// carries for the overlay interface. These tests create their own utun device, and
// stdnet's filter probes with wgctrl every interface it is not told to skip, which
// on a userspace WireGuard platform reaches the UAPI socket of this same process.
// Declared here rather than imported because profilemanager imports this package.
var testIFaceBlackList = []string{"wt", "utun", "tun0"}
func TestWGIface_UpdateAddr(t *testing.T) { func TestWGIface_UpdateAddr(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
addr := "100.64.0.1/8" addr := "100.64.0.1/8"
wgPort := 33100 wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
if err != nil {
t.Fatal(err)
}
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
@@ -127,10 +131,7 @@ func getIfaceAddrs(ifaceName string) ([]net.Addr, error) {
func Test_CreateInterface(t *testing.T) { func Test_CreateInterface(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1)
wgIP := "10.99.99.1/32" wgIP := "10.99.99.1/32"
newNet, err := stdnet.NewNet(context.Background(), nil) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
if err != nil {
t.Fatal(err)
}
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
Address: wgaddr.MustParseWGAddress(wgIP), Address: wgaddr.MustParseWGAddress(wgIP),
@@ -170,10 +171,7 @@ func Test_Close(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
wgIP := "10.99.99.2/32" wgIP := "10.99.99.2/32"
wgPort := 33100 wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
if err != nil {
t.Fatal(err)
}
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
@@ -215,10 +213,7 @@ func TestRecreation(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
wgIP := "10.99.99.2/32" wgIP := "10.99.99.2/32"
wgPort := 33100 wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
if err != nil {
t.Fatal(err)
}
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
@@ -288,10 +283,7 @@ func Test_ConfigureInterface(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3)
wgIP := "10.99.99.5/30" wgIP := "10.99.99.5/30"
wgPort := 33100 wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
if err != nil {
t.Fatal(err)
}
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
Address: wgaddr.MustParseWGAddress(wgIP), Address: wgaddr.MustParseWGAddress(wgIP),
@@ -343,10 +335,7 @@ func Test_ConfigureInterface(t *testing.T) {
func Test_UpdatePeer(t *testing.T) { func Test_UpdatePeer(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
wgIP := "10.99.99.9/30" wgIP := "10.99.99.9/30"
newNet, err := stdnet.NewNet(context.Background(), nil) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
if err != nil {
t.Fatal(err)
}
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
@@ -413,10 +402,7 @@ func Test_UpdatePeer(t *testing.T) {
func Test_RemovePeer(t *testing.T) { func Test_RemovePeer(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
wgIP := "10.99.99.13/30" wgIP := "10.99.99.13/30"
newNet, err := stdnet.NewNet(context.Background(), nil) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
if err != nil {
t.Fatal(err)
}
opts := WGIFaceOpts{ opts := WGIFaceOpts{
IFaceName: ifaceName, IFaceName: ifaceName,
@@ -477,10 +463,7 @@ func Test_ConnectPeers(t *testing.T) {
peer2wgPort := 33200 peer2wgPort := 33200
keepAlive := 1 * time.Second keepAlive := 1 * time.Second
newNet, err := stdnet.NewNet(context.Background(), nil) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
if err != nil {
t.Fatal(err)
}
guid := fmt.Sprintf("{%s}", uuid.New().String()) guid := fmt.Sprintf("{%s}", uuid.New().String())
device.CustomWindowsGUIDString = strings.ToLower(guid) device.CustomWindowsGUIDString = strings.ToLower(guid)
@@ -516,10 +499,7 @@ func Test_ConnectPeers(t *testing.T) {
guid = fmt.Sprintf("{%s}", uuid.New().String()) guid = fmt.Sprintf("{%s}", uuid.New().String())
device.CustomWindowsGUIDString = strings.ToLower(guid) device.CustomWindowsGUIDString = strings.ToLower(guid)
newNet, err = stdnet.NewNet(context.Background(), nil) newNet = stdnet.NewNet(context.Background(), testIFaceBlackList)
if err != nil {
t.Fatal(err)
}
optsPeer2 := WGIFaceOpts{ optsPeer2 := WGIFaceOpts{
IFaceName: peer2ifaceName, IFaceName: peer2ifaceName,
@@ -568,11 +548,14 @@ func Test_ConnectPeers(t *testing.T) {
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
// The peers use userspace WireGuard (stdnet transport). A tight busy-loop // On Linux with the kernel module both peers are kernel devices, elsewhere
// here starves the wireguard-go goroutines that process the handshake, so // they run on wireguard-go. A tight busy-loop here would starve the
// poll on a ticker instead and yield the CPU between checks. WireGuard also // wireguard-go goroutines that process the handshake, so poll on a ticker
// only retries a lost handshake initiation every REKEY_TIMEOUT (5s), which // instead and yield the CPU between checks. WireGuard also only retries a
// is why the overall wait can occasionally stretch to tens of seconds. // lost handshake initiation every REKEY_TIMEOUT (5s), which is why the
// overall wait can occasionally stretch to tens of seconds. Each side sends
// its first initiation when its peer is configured, and the first one leaves
// before the other device knows the peer, so that one is always wasted.
timeout := 30 * time.Second timeout := 30 * time.Second
timeoutChannel := time.After(timeout) timeoutChannel := time.After(timeout)
ticker := time.NewTicker(500 * time.Millisecond) ticker := time.NewTicker(500 * time.Millisecond)
@@ -590,13 +573,26 @@ func Test_ConnectPeers(t *testing.T) {
select { select {
case <-timeoutChannel: case <-timeoutChannel:
t.Fatalf("waiting for peer handshake timeout after %s", timeout.String()) // The counters tell whether initiations were sent at all, whether they
// arrived, and whether only one direction is working.
t.Fatalf("waiting for peer handshake timeout after %s\n%s\n%s", timeout.String(),
describePeer(peer1ifaceName, peer2Key.PublicKey().String()),
describePeer(peer2ifaceName, peer1Key.PublicKey().String()))
case <-ticker.C: case <-ticker.C:
} }
} }
} }
func describePeer(ifaceName, peerPubKey string) string {
peer, err := getPeer(ifaceName, peerPubKey)
if err != nil {
return fmt.Sprintf("%s: peer %s: %v", ifaceName, peerPubKey, err)
}
return fmt.Sprintf("%s: peer %s endpoint=%v tx=%d rx=%d last_handshake=%v",
ifaceName, peerPubKey, peer.Endpoint, peer.TransmitBytes, peer.ReceiveBytes, peer.LastHandshakeTime)
}
func getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) { func getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) {
wg, err := wgctrl.New() wg, err := wgctrl.New()
if err != nil { if err != nil {
+1 -4
View File
@@ -200,10 +200,7 @@ func (m *SingleSocketUDPMux) updateLocalAddresses() {
} }
if len(networks) > 0 { if len(networks) > 0 {
if m.params.Net == nil { if m.params.Net == nil {
var err error m.params.Net = stdnet.NewNet(context.Background(), nil)
if m.params.Net, err = stdnet.NewNet(context.Background(), nil); err != nil {
m.params.Logger.Errorf("failed to get create network: %v", err)
}
} }
ips, err := localInterfaces(m.params.Net, m.params.InterfaceFilter, nil, networks, true) ips, err := localInterfaces(m.params.Net, m.params.InterfaceFilter, nil, networks, true)
-32
View File
@@ -1,32 +0,0 @@
package ebpf
import (
"fmt"
"net"
)
var (
portRangeStart = 3128
portRangeEnd = portRangeStart + 100
)
type portLookup struct {
}
func (pl portLookup) searchFreePort() (int, error) {
for i := portRangeStart; i <= portRangeEnd; i++ {
if pl.tryToBind(i) == nil {
return i, nil
}
}
return 0, fmt.Errorf("failed to bind free port for eBPF proxy")
}
func (pl portLookup) tryToBind(port int) error {
l, err := net.ListenPacket("udp", fmt.Sprintf(":%d", port))
if err != nil {
return err
}
_ = l.Close()
return nil
}
@@ -1,45 +0,0 @@
package ebpf
import (
"fmt"
"net"
"testing"
)
func Test_portLookup_searchFreePort(t *testing.T) {
pl := portLookup{}
_, err := pl.searchFreePort()
if err != nil {
t.Fatal(err)
}
}
func Test_portLookup_on_allocated(t *testing.T) {
pl := portLookup{}
portRangeStart = 4128
portRangeEnd = portRangeStart + 100
allocatedPort, err := allocatePort(portRangeStart)
if err != nil {
t.Fatal(err)
}
defer allocatedPort.Close()
fp, err := pl.searchFreePort()
if err != nil {
t.Fatal(err)
}
if fp != (portRangeStart + 1) {
t.Errorf("invalid free port, expected: %d, got: %d", portRangeStart+1, fp)
}
}
func allocatePort(port int) (net.PacketConn, error) {
c, err := net.ListenPacket("udp", fmt.Sprintf(":%d", port))
if err != nil {
return nil, err
}
return c, err
}
-243
View File
@@ -1,243 +0,0 @@
//go:build linux && !android
package ebpf
import (
"context"
"fmt"
"net"
"sync"
"github.com/hashicorp/go-multierror"
"github.com/pion/transport/v3"
log "github.com/sirupsen/logrus"
nberrors "github.com/netbirdio/netbird/client/errors"
"github.com/netbirdio/netbird/client/iface/bufsize"
"github.com/netbirdio/netbird/client/iface/wgproxy/rawsocket"
"github.com/netbirdio/netbird/client/internal/ebpf"
ebpfMgr "github.com/netbirdio/netbird/client/internal/ebpf/manager"
nbnet "github.com/netbirdio/netbird/client/net"
)
const (
loopbackAddr = "127.0.0.1"
)
// WGEBPFProxy definition for proxy with EBPF support
type WGEBPFProxy struct {
localWGListenPort int
proxyPort int
mtu uint16
ebpfManager ebpfMgr.Manager
relayedConnStore map[uint16]net.Conn
relayedConnMutex sync.Mutex
lastUsedPort uint16
rawConnIPv4 net.PacketConn
rawConnIPv6 net.PacketConn
conn transport.UDPConn
ctx context.Context
ctxCancel context.CancelFunc
}
// NewWGEBPFProxy create new WGEBPFProxy instance
func NewWGEBPFProxy(wgPort int, mtu uint16) *WGEBPFProxy {
log.Debugf("instantiate ebpf proxy")
wgProxy := &WGEBPFProxy{
localWGListenPort: wgPort,
mtu: mtu,
ebpfManager: ebpf.GetEbpfManagerInstance(),
relayedConnStore: make(map[uint16]net.Conn),
}
return wgProxy
}
// Listen load ebpf program and listen the proxy
func (p *WGEBPFProxy) Listen() error {
pl := portLookup{}
proxyPort, err := pl.searchFreePort()
if err != nil {
return err
}
p.proxyPort = proxyPort
// Prepare IPv4 raw socket (required)
p.rawConnIPv4, err = rawsocket.PrepareSenderRawSocketIPv4()
if err != nil {
return err
}
// Prepare IPv6 raw socket (optional)
p.rawConnIPv6, err = rawsocket.PrepareSenderRawSocketIPv6()
if err != nil {
log.Warnf("failed to prepare IPv6 raw socket, continuing with IPv4 only: %v", err)
}
err = p.ebpfManager.LoadWgProxy(proxyPort, p.localWGListenPort)
if err != nil {
if closeErr := p.rawConnIPv4.Close(); closeErr != nil {
log.Warnf("failed to close IPv4 raw socket: %v", closeErr)
}
if p.rawConnIPv6 != nil {
if closeErr := p.rawConnIPv6.Close(); closeErr != nil {
log.Warnf("failed to close IPv6 raw socket: %v", closeErr)
}
}
return err
}
addr := net.UDPAddr{
Port: proxyPort,
IP: net.ParseIP(loopbackAddr),
}
p.ctx, p.ctxCancel = context.WithCancel(context.Background())
conn, err := nbnet.ListenUDP("udp", &addr)
if err != nil {
if cErr := p.Free(); cErr != nil {
log.Errorf("Failed to close the wgproxy: %s", cErr)
}
return err
}
p.conn = conn
go p.proxyToRemote()
log.Infof("local wg proxy listening on: %d", proxyPort)
return nil
}
// AddRelayedConn add new relayed connection for the proxy
func (p *WGEBPFProxy) AddRelayedConn(relayedConn net.Conn) (*net.UDPAddr, error) {
wgEndpointPort, err := p.storeRelayedConn(relayedConn)
if err != nil {
return nil, err
}
log.Infof("relayed conn added to wg proxy store: %s, endpoint port: :%d", relayedConn.RemoteAddr(), wgEndpointPort)
wgEndpoint := &net.UDPAddr{
IP: net.ParseIP(loopbackAddr),
Port: int(wgEndpointPort),
}
return wgEndpoint, nil
}
// Free resources except the remoteConns will be keep open.
func (p *WGEBPFProxy) Free() error {
log.Debugf("free up ebpf wg proxy")
if p.ctx != nil && p.ctx.Err() != nil {
//nolint
return nil
}
p.ctxCancel()
var result *multierror.Error
if p.conn != nil {
if err := p.conn.Close(); err != nil {
result = multierror.Append(result, err)
}
}
if err := p.ebpfManager.FreeWGProxy(); err != nil {
result = multierror.Append(result, err)
}
if p.rawConnIPv4 != nil {
if err := p.rawConnIPv4.Close(); err != nil {
result = multierror.Append(result, err)
}
}
if p.rawConnIPv6 != nil {
if err := p.rawConnIPv6.Close(); err != nil {
result = multierror.Append(result, err)
}
}
return nberrors.FormatErrorOrNil(result)
}
// GetProxyPort returns the proxy listening port.
func (p *WGEBPFProxy) GetProxyPort() uint16 {
return uint16(p.proxyPort)
}
// proxyToRemote read messages from local WireGuard interface and forward it to remote conn
// From this go routine has only one instance.
func (p *WGEBPFProxy) proxyToRemote() {
buf := make([]byte, p.mtu+bufsize.WGBufferOverhead)
for p.ctx.Err() == nil {
if err := p.readAndForwardPacket(buf); err != nil {
if p.ctx.Err() != nil {
return
}
log.Errorf("failed to proxy packet to remote conn: %s", err)
}
}
}
func (p *WGEBPFProxy) readAndForwardPacket(buf []byte) error {
n, addr, err := p.conn.ReadFromUDP(buf)
if err != nil {
return fmt.Errorf("failed to read UDP packet from WG: %w", err)
}
p.relayedConnMutex.Lock()
conn, ok := p.relayedConnStore[uint16(addr.Port)]
p.relayedConnMutex.Unlock()
if !ok {
if p.ctx.Err() == nil {
log.Debugf("relayed conn not found by port because conn already has been closed: %d", addr.Port)
}
return nil
}
if _, err := conn.Write(buf[:n]); err != nil {
return fmt.Errorf("forward local WG packet (%d) to remote relayed conn: %w", addr.Port, err)
}
return nil
}
func (p *WGEBPFProxy) storeRelayedConn(relayedConn net.Conn) (uint16, error) {
p.relayedConnMutex.Lock()
defer p.relayedConnMutex.Unlock()
np, err := p.nextFreePort()
if err != nil {
return np, err
}
p.relayedConnStore[np] = relayedConn
return np, nil
}
func (p *WGEBPFProxy) removeRelayedConn(relayedConnID uint16) {
p.relayedConnMutex.Lock()
defer p.relayedConnMutex.Unlock()
_, ok := p.relayedConnStore[relayedConnID]
if ok {
log.Debugf("remove relayed conn from store by port: %d", relayedConnID)
}
delete(p.relayedConnStore, relayedConnID)
}
func (p *WGEBPFProxy) nextFreePort() (uint16, error) {
if len(p.relayedConnStore) == 65535 {
return 0, fmt.Errorf("reached maximum relayed connection numbers")
}
generatePort:
if p.lastUsedPort == 65535 {
p.lastUsedPort = 1
} else {
p.lastUsedPort++
}
if _, ok := p.relayedConnStore[p.lastUsedPort]; ok {
goto generatePort
}
return p.lastUsedPort, nil
}
-56
View File
@@ -1,56 +0,0 @@
//go:build linux && !android
package ebpf
import (
"testing"
)
func TestWGEBPFProxy_connStore(t *testing.T) {
wgProxy := NewWGEBPFProxy(1, 1280)
p, _ := wgProxy.storeRelayedConn(nil)
if p != 1 {
t.Errorf("invalid initial port: %d", wgProxy.lastUsedPort)
}
numOfConns := 10
for i := 0; i < numOfConns; i++ {
p, _ = wgProxy.storeRelayedConn(nil)
}
if p != uint16(numOfConns)+1 {
t.Errorf("invalid last used port: %d, expected: %d", p, numOfConns+1)
}
if len(wgProxy.relayedConnStore) != numOfConns+1 {
t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.relayedConnStore), numOfConns+1)
}
}
func TestWGEBPFProxy_portCalculation_overflow(t *testing.T) {
wgProxy := NewWGEBPFProxy(1, 1280)
_, _ = wgProxy.storeRelayedConn(nil)
wgProxy.lastUsedPort = 65535
p, _ := wgProxy.storeRelayedConn(nil)
if len(wgProxy.relayedConnStore) != 2 {
t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.relayedConnStore), 2)
}
if p != 2 {
t.Errorf("invalid last used port: %d, expected: %d", p, 2)
}
}
func TestWGEBPFProxy_portCalculation_maxConn(t *testing.T) {
wgProxy := NewWGEBPFProxy(1, 1280)
for i := 0; i < 65535; i++ {
_, _ = wgProxy.storeRelayedConn(nil)
}
_, err := wgProxy.storeRelayedConn(nil)
if err == nil {
t.Errorf("invalid relayed conn store calculation")
}
}
+27 -24
View File
@@ -8,11 +8,13 @@ import (
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/iface/wgproxy/ebpf" "github.com/netbirdio/netbird/client/iface/wgproxy/loopback"
udpProxy "github.com/netbirdio/netbird/client/iface/wgproxy/udp" udpProxy "github.com/netbirdio/netbird/client/iface/wgproxy/udp"
) )
const ( const (
envDisableKernelWGProxy = "NB_DISABLE_KERNEL_WG_PROXY"
// envDisableEBPFWGProxy is a deprecated alias for envDisableKernelWGProxy.
envDisableEBPFWGProxy = "NB_DISABLE_EBPF_WG_PROXY" envDisableEBPFWGProxy = "NB_DISABLE_EBPF_WG_PROXY"
) )
@@ -20,7 +22,7 @@ type KernelFactory struct {
wgPort int wgPort int
mtu uint16 mtu uint16
ebpfProxy *ebpf.WGEBPFProxy loopbackProxy *loopback.Proxy
} }
func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory { func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory {
@@ -29,55 +31,56 @@ func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory {
mtu: mtu, mtu: mtu,
} }
if isEBPFDisabled() { if isKernelProxyDisabled() {
log.Infof("WireGuard Proxy Factory will produce UDP proxy") log.Infof("WireGuard Proxy Factory will produce UDP proxy")
log.Infof("eBPF WireGuard proxy is disabled via %s environment variable", envDisableEBPFWGProxy)
return f return f
} }
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, mtu) loopbackProxy := loopback.NewProxy(wgPort, mtu)
if err := ebpfProxy.Listen(); err != nil { if err := loopbackProxy.Listen(); err != nil {
log.Infof("WireGuard Proxy Factory will produce UDP proxy") log.Infof("WireGuard Proxy Factory will produce UDP proxy")
log.Warnf("failed to initialize ebpf proxy, fallback to user space proxy: %s", err) log.Warnf("failed to initialize loopback proxy, fallback to user space proxy: %s", err)
return f return f
} }
log.Infof("WireGuard Proxy Factory will produce eBPF proxy") log.Infof("WireGuard Proxy Factory will produce loopback proxy")
f.ebpfProxy = ebpfProxy f.loopbackProxy = loopbackProxy
return f return f
} }
func (w *KernelFactory) GetProxy() Proxy { func (w *KernelFactory) GetProxy() Proxy {
if w.ebpfProxy == nil { if w.loopbackProxy == nil {
return udpProxy.NewWGUDPProxy(w.wgPort, w.mtu) return udpProxy.NewWGUDPProxy(w.wgPort, w.mtu)
} }
return ebpf.NewProxyWrapper(w.ebpfProxy) return loopback.NewProxyWrapper(w.loopbackProxy)
}
// GetProxyPort returns the eBPF proxy port, or 0 if eBPF is not active.
func (w *KernelFactory) GetProxyPort() uint16 {
if w.ebpfProxy == nil {
return 0
}
return w.ebpfProxy.GetProxyPort()
} }
func (w *KernelFactory) Free() error { func (w *KernelFactory) Free() error {
if w.ebpfProxy == nil { if w.loopbackProxy == nil {
return nil return nil
} }
return w.ebpfProxy.Free() return w.loopbackProxy.Free()
} }
func isEBPFDisabled() bool { func isKernelProxyDisabled() bool {
val := os.Getenv(envDisableEBPFWGProxy) env := envDisableKernelWGProxy
val := os.Getenv(env)
if val == "" {
env = envDisableEBPFWGProxy
val = os.Getenv(env)
}
if val == "" { if val == "" {
return false return false
} }
disabled, err := strconv.ParseBool(val) disabled, err := strconv.ParseBool(val)
if err != nil { if err != nil {
log.Warnf("failed to parse %s: %v", envDisableEBPFWGProxy, err) log.Warnf("failed to parse %s: %v", env, err)
return false return false
} }
if disabled {
log.Infof("kernel WireGuard proxy is disabled via %s", env)
}
return disabled return disabled
} }
-5
View File
@@ -24,11 +24,6 @@ func (w *USPFactory) GetProxy() Proxy {
return proxyBind.NewProxyBind(w.bind, w.mtu) return proxyBind.NewProxyBind(w.bind, w.mtu)
} }
// GetProxyPort returns 0 as userspace WireGuard doesn't use a separate proxy port.
func (w *USPFactory) GetProxyPort() uint16 {
return 0
}
func (w *USPFactory) Free() error { func (w *USPFactory) Free() error {
return nil return nil
} }
+70
View File
@@ -0,0 +1,70 @@
//go:build linux && !android
package loopback
import (
"fmt"
"net/netip"
)
// Peer endpoints live in the upper half of 127.0.0.0/8. Everything in that
// range is delivered to the loopback device without any address or route being
// configured, and staying out of 127.0.0.0/9 keeps well-known squatters such as
// 127.0.0.53 (systemd-resolved) and 127.0.1.1 out of the way.
const (
addrRangeBase uint32 = 0x7f800000 // 127.128.0.0
addrRangeSize uint32 = 1 << 23 // /9
addrRangePrefix = "127.128.0.0/9"
)
// allocator hands out one loopback address per relayed connection. The address
// is the peer's identity: WireGuard sends to it, and the proxy recovers which
// peer a packet belongs to from the destination address.
type allocator struct {
cursor uint32
}
// next returns the first free address at or after the cursor, wrapping once.
// inUse reports whether an address is already handed out.
func (a *allocator) next(inUse func(netip.Addr) bool) (netip.Addr, error) {
for i := uint32(0); i < addrRangeSize; i++ {
a.cursor = (a.cursor + 1) % addrRangeSize
addr := addrFromOffset(a.cursor)
if !addr.IsValid() {
continue
}
if inUse(addr) {
continue
}
return addr, nil
}
return netip.Addr{}, fmt.Errorf("no free endpoint address in %s", addrRangePrefix)
}
// addrFromOffset maps an offset in the range to an address, skipping the .0 and
// .255 hosts. They are unremarkable on loopback, but tools and firewall rules
// tend to treat them as network and broadcast addresses.
func addrFromOffset(offset uint32) netip.Addr {
last := offset & 0xff
if last == 0 || last == 0xff {
return netip.Addr{}
}
v := addrRangeBase + offset
return netip.AddrFrom4([4]byte{
byte(v >> 24),
byte(v >> 16),
byte(v >> 8),
byte(v),
})
}
// inRange reports whether addr is one this proxy could have handed out.
func inRange(addr netip.Addr) bool {
if !addr.Is4() {
return false
}
b := addr.As4()
v := uint32(b[0])<<24 | uint32(b[1])<<16 | uint32(b[2])<<8 | uint32(b[3])
return v >= addrRangeBase && v < addrRangeBase+addrRangeSize && b[3] != 0 && b[3] != 0xff
}
+114
View File
@@ -0,0 +1,114 @@
//go:build linux && !android
package loopback
import (
"net/netip"
"testing"
)
func TestAllocatorHandsOutDistinctAddresses(t *testing.T) {
var a allocator
taken := make(map[netip.Addr]bool)
for i := 0; i < 1000; i++ {
addr, err := a.next(func(candidate netip.Addr) bool { return taken[candidate] })
if err != nil {
t.Fatalf("allocate %d: %v", i, err)
}
if taken[addr] {
t.Fatalf("address %s handed out twice", addr)
}
if !inRange(addr) {
t.Fatalf("address %s outside %s", addr, addrRangePrefix)
}
taken[addr] = true
}
}
func TestAllocatorSkipsNetworkAndBroadcastHosts(t *testing.T) {
var a allocator
taken := make(map[netip.Addr]bool)
// enough allocations to walk past a .255/.0 boundary
for i := 0; i < 600; i++ {
addr, err := a.next(func(candidate netip.Addr) bool { return taken[candidate] })
if err != nil {
t.Fatalf("allocate %d: %v", i, err)
}
last := addr.As4()[3]
if last == 0 || last == 255 {
t.Fatalf("address %s ends in .%d", addr, last)
}
taken[addr] = true
}
}
func TestAllocatorReusesReleasedAddresses(t *testing.T) {
var a allocator
taken := make(map[netip.Addr]bool)
inUse := func(candidate netip.Addr) bool { return taken[candidate] }
alloc := func() netip.Addr {
t.Helper()
addr, err := a.next(inUse)
if err != nil {
t.Fatalf("allocate: %v", err)
}
taken[addr] = true
return addr
}
first := alloc()
second := alloc()
delete(taken, first)
// The cursor only moves forward, so a released address comes back after a
// wrap. Park the cursor near the end of the range instead of allocating
// 2^23 addresses: the next call takes the last usable address, and the one
// after that wraps past the skipped .255 and .0 hosts to the released one.
a.cursor = addrRangeSize - 3
last := alloc()
if want := netip.MustParseAddr("127.255.255.254"); last != want {
t.Fatalf("expected the last usable address %s before the wrap, got %s", want, last)
}
if reused := alloc(); reused != first {
t.Fatalf("expected the released address %s after the wrap, got %s", first, reused)
}
// second is still held, so the allocator must step over it.
if next := alloc(); next == second {
t.Fatalf("allocator handed out %s while it was still in use", second)
}
}
func TestInRange(t *testing.T) {
tests := []struct {
addr string
want bool
}{
{"127.128.0.1", true},
{"127.255.255.254", true},
{"127.128.0.0", false}, // network host, never handed out
{"127.128.5.255", false}, // broadcast host, never handed out
{"127.127.255.255", false}, // below the range, where 127.0.0.53 and friends live
{"127.0.0.1", false},
{"127.0.0.53", false},
{"127.0.1.1", false},
{"128.0.0.1", false},
{"10.0.0.1", false},
}
for _, tc := range tests {
addr := netip.MustParseAddr(tc.addr)
if got := inRange(addr); got != tc.want {
t.Errorf("inRange(%s) = %v, want %v", tc.addr, got, tc.want)
}
}
}
func TestInRangeIgnoresIPv6(t *testing.T) {
if inRange(netip.MustParseAddr("::1")) {
t.Error("inRange(::1) = true, want false")
}
}
+291
View File
@@ -0,0 +1,291 @@
//go:build linux && !android
package loopback
import (
"context"
"fmt"
"net"
"net/netip"
"sync"
"syscall"
"github.com/hashicorp/go-multierror"
log "github.com/sirupsen/logrus"
"golang.org/x/net/ipv4"
"golang.org/x/sys/unix"
nberrors "github.com/netbirdio/netbird/client/errors"
"github.com/netbirdio/netbird/client/iface/bufsize"
"github.com/netbirdio/netbird/client/iface/wgproxy/rawsocket"
)
const (
loopbackDevice = "lo"
portRangeStart = 3128
portRangeEnd = portRangeStart + 100
)
// Proxy forwards packets between relayed connections and a local kernel
// WireGuard instance. Every relayed peer gets its own loopback address as its
// WireGuard endpoint, so a single socket serves all of them: the destination
// address of an incoming packet identifies the peer.
type Proxy struct {
localWGListenPort int
mtu uint16
proxyPort int
conn *net.UDPConn
packetConn *ipv4.PacketConn
loIndex int
rawConnIPv4 net.PacketConn
rawConnIPv6 net.PacketConn
relayedConnMutex sync.Mutex
relayedConnStore map[netip.Addr]net.Conn
addrs allocator
ctx context.Context
ctxCancel context.CancelFunc
}
// NewProxy creates a proxy for the WireGuard instance listening on wgPort.
func NewProxy(wgPort int, mtu uint16) *Proxy {
log.Debugf("instantiate loopback wg proxy")
return &Proxy{
localWGListenPort: wgPort,
mtu: mtu,
relayedConnStore: make(map[netip.Addr]net.Conn),
}
}
// Listen opens the shared socket and starts forwarding WireGuard packets to the
// relayed connections.
func (p *Proxy) Listen() error {
rawConnIPv4, err := rawsocket.PrepareSenderRawSocketIPv4()
if err != nil {
return fmt.Errorf("prepare IPv4 raw socket: %w", err)
}
p.rawConnIPv4 = rawConnIPv4
p.rawConnIPv6, err = rawsocket.PrepareSenderRawSocketIPv6()
if err != nil {
log.Warnf("failed to prepare IPv6 raw socket, continuing with IPv4 only: %v", err)
}
loopback, err := net.InterfaceByName(loopbackDevice)
if err != nil {
if freeErr := p.Free(); freeErr != nil {
log.Errorf("failed to free the wgproxy: %s", freeErr)
}
return fmt.Errorf("look up %s: %w", loopbackDevice, err)
}
p.loIndex = loopback.Index
if err := p.listen(); err != nil {
if freeErr := p.Free(); freeErr != nil {
log.Errorf("failed to free the wgproxy: %s", freeErr)
}
return err
}
p.ctx, p.ctxCancel = context.WithCancel(context.Background())
go p.proxyToRemote()
log.Infof("local wg proxy listening on %s:%d", addrRangePrefix, p.proxyPort)
return nil
}
// listen binds the shared socket on the first free port of the range. The bind
// has to be a wildcard one to receive every peer address in the range, so it is
// restricted to the loopback device: without that the port would be reachable
// on every interface.
func (p *Proxy) listen() error {
var lastErr error
for port := portRangeStart; port <= portRangeEnd; port++ {
err := p.listenOn(port)
if err == nil {
p.proxyPort = port
return nil
}
lastErr = err
}
return fmt.Errorf("bind proxy port in range %d-%d: %w", portRangeStart, portRangeEnd, lastErr)
}
func (p *Proxy) listenOn(proxyPort int) error {
lc := net.ListenConfig{
Control: func(_, _ string, c syscall.RawConn) error {
var sockErr error
if err := c.Control(func(fd uintptr) {
if err := unix.SetsockoptString(int(fd), unix.SOL_SOCKET, unix.SO_BINDTODEVICE, loopbackDevice); err != nil {
sockErr = fmt.Errorf("bind to %s: %w", loopbackDevice, err)
return
}
}); err != nil {
return fmt.Errorf("control socket: %w", err)
}
return sockErr
},
}
conn, err := lc.ListenPacket(context.Background(), "udp4", fmt.Sprintf(":%d", proxyPort))
if err != nil {
return fmt.Errorf("listen on :%d: %w", proxyPort, err)
}
udpConn, ok := conn.(*net.UDPConn)
if !ok {
if closeErr := conn.Close(); closeErr != nil {
log.Errorf("failed to close proxy conn: %s", closeErr)
}
return fmt.Errorf("unexpected conn type %T", conn)
}
packetConn := ipv4.NewPacketConn(udpConn)
// the destination address carries the peer identity, the interface index is
// checked on receive as a second line of defense behind SO_BINDTODEVICE
if err := packetConn.SetControlMessage(ipv4.FlagDst|ipv4.FlagInterface, true); err != nil {
if closeErr := udpConn.Close(); closeErr != nil {
log.Errorf("failed to close proxy conn: %s", closeErr)
}
return fmt.Errorf("request destination address: %w", err)
}
p.conn = udpConn
p.packetConn = packetConn
return nil
}
// AddRelayedConn assigns an endpoint address to the relayed connection and
// returns the address WireGuard should send to, along with the key the
// connection is stored under.
func (p *Proxy) AddRelayedConn(relayedConn net.Conn) (*net.UDPAddr, netip.Addr, error) {
addr, err := p.storeRelayedConn(relayedConn)
if err != nil {
return nil, netip.Addr{}, err
}
log.Infof("relayed conn added to wg proxy store: %s, endpoint address: %s", relayedConn.RemoteAddr(), addr)
return &net.UDPAddr{
IP: addr.AsSlice(),
Port: p.proxyPort,
}, addr, nil
}
// Free releases the proxy resources. The relayed connections are left open.
func (p *Proxy) Free() error {
log.Debugf("free up loopback wg proxy")
if p.ctx != nil && p.ctx.Err() != nil {
//nolint
return nil
}
if p.ctxCancel != nil {
p.ctxCancel()
}
var result *multierror.Error
if p.conn != nil {
if err := p.conn.Close(); err != nil {
result = multierror.Append(result, err)
}
}
if p.rawConnIPv4 != nil {
if err := p.rawConnIPv4.Close(); err != nil {
result = multierror.Append(result, err)
}
}
if p.rawConnIPv6 != nil {
if err := p.rawConnIPv6.Close(); err != nil {
result = multierror.Append(result, err)
}
}
return nberrors.FormatErrorOrNil(result)
}
// proxyToRemote reads packets from the local WireGuard instance and forwards
// them to the relayed connection the destination address belongs to.
func (p *Proxy) proxyToRemote() {
buf := make([]byte, p.mtu+bufsize.WGBufferOverhead)
for p.ctx.Err() == nil {
if err := p.readAndForwardPacket(buf); err != nil {
if p.ctx.Err() != nil {
return
}
log.Errorf("failed to proxy packet to remote conn: %s", err)
}
}
}
func (p *Proxy) readAndForwardPacket(buf []byte) error {
n, cm, _, err := p.packetConn.ReadFrom(buf)
if err != nil {
return fmt.Errorf("read UDP packet from WG: %w", err)
}
if cm == nil {
return fmt.Errorf("no control message on packet")
}
if cm.IfIndex != p.loIndex {
log.Tracef("dropping packet received on interface %d instead of %s", cm.IfIndex, loopbackDevice)
return nil
}
dst, ok := netip.AddrFromSlice(cm.Dst.To4())
if !ok || !inRange(dst) {
log.Tracef("dropping packet for unexpected destination %s", cm.Dst)
return nil
}
p.relayedConnMutex.Lock()
conn, ok := p.relayedConnStore[dst]
p.relayedConnMutex.Unlock()
if !ok {
if p.ctx.Err() == nil {
log.Debugf("relayed conn not found by address because conn already has been closed: %s", dst)
}
return nil
}
if _, err := conn.Write(buf[:n]); err != nil {
return fmt.Errorf("forward local WG packet (%s) to remote relayed conn: %w", dst, err)
}
return nil
}
func (p *Proxy) storeRelayedConn(relayedConn net.Conn) (netip.Addr, error) {
p.relayedConnMutex.Lock()
defer p.relayedConnMutex.Unlock()
addr, err := p.addrs.next(func(a netip.Addr) bool {
_, ok := p.relayedConnStore[a]
return ok
})
if err != nil {
return netip.Addr{}, err
}
p.relayedConnStore[addr] = relayedConn
return addr, nil
}
// removeRelayedConn releases an endpoint address. It only removes the entry
// while it still belongs to relayedConn, so a late release cannot take an
// address away from the peer it was handed to next.
func (p *Proxy) removeRelayedConn(addr netip.Addr, relayedConn net.Conn) {
p.relayedConnMutex.Lock()
defer p.relayedConnMutex.Unlock()
if stored, ok := p.relayedConnStore[addr]; !ok || stored != relayedConn {
return
}
log.Debugf("remove relayed conn from store by address: %s", addr)
delete(p.relayedConnStore, addr)
}
@@ -0,0 +1,196 @@
//go:build linux && !android && privileged
package loopback
import (
"context"
"net"
"strconv"
"testing"
"time"
)
const testWGPort = 51862
// relayEnd stands in for a relayed connection: the proxy writes what it read
// from WireGuard into it, and the test reads it back out here.
func relayEnd(t *testing.T) (proxySide net.Conn, testSide *net.UDPConn) {
t.Helper()
testSide, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatalf("relay listener: %v", err)
}
t.Cleanup(func() {
if err := testSide.Close(); err != nil {
t.Logf("close relay listener: %v", err)
}
})
proxySide, err = net.Dial("udp", testSide.LocalAddr().String())
if err != nil {
t.Fatalf("relay conn: %v", err)
}
t.Cleanup(func() {
if err := proxySide.Close(); err != nil {
t.Logf("close relay conn: %v", err)
}
})
return proxySide, testSide
}
// TestProxyDemuxesByDestinationAddress is the core of the design: one socket
// serves every peer, and the destination address decides which relayed
// connection a WireGuard packet belongs to.
func TestProxyDemuxesByDestinationAddress(t *testing.T) {
proxy := NewProxy(testWGPort, 1280)
if err := proxy.Listen(); err != nil {
t.Fatalf("listen: %v", err)
}
defer func() {
if err := proxy.Free(); err != nil {
t.Errorf("free proxy: %v", err)
}
}()
const peers = 3
endpoints := make([]*net.UDPAddr, 0, peers)
readers := make([]*net.UDPConn, 0, peers)
for i := 0; i < peers; i++ {
proxySide, testSide := relayEnd(t)
endpoint, _, err := proxy.AddRelayedConn(proxySide)
if err != nil {
t.Fatalf("add relayed conn %d: %v", i, err)
}
if endpoint.Port != proxy.proxyPort {
t.Errorf("peer %d endpoint port = %d, want the shared proxy port %d", i, endpoint.Port, proxy.proxyPort)
}
endpoints = append(endpoints, endpoint)
readers = append(readers, testSide)
}
// every peer must have its own address, otherwise they are indistinguishable
seen := make(map[string]bool, peers)
for i, endpoint := range endpoints {
if seen[endpoint.IP.String()] {
t.Fatalf("peer %d reuses endpoint address %s", i, endpoint.IP)
}
seen[endpoint.IP.String()] = true
}
wgSock, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatalf("wg socket: %v", err)
}
defer func() {
if err := wgSock.Close(); err != nil {
t.Logf("close wg socket: %v", err)
}
}()
for i, endpoint := range endpoints {
payload := []byte{byte(i), 'p', 'k', 't'}
if _, err := wgSock.WriteTo(payload, endpoint); err != nil {
t.Fatalf("write to peer %d endpoint %s: %v", i, endpoint, err)
}
buf := make([]byte, 1500)
if err := readers[i].SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil {
t.Fatalf("set read deadline: %v", err)
}
n, _, err := readers[i].ReadFrom(buf)
if err != nil {
t.Fatalf("peer %d did not receive its packet: %v", i, err)
}
if string(buf[:n]) != string(payload) {
t.Errorf("peer %d got %q, want %q", i, buf[:n], payload)
}
// no other peer may see it
for j, other := range readers {
if j == i {
continue
}
if err := other.SetReadDeadline(time.Now().Add(100 * time.Millisecond)); err != nil {
t.Fatalf("set read deadline: %v", err)
}
if _, _, err := other.ReadFrom(buf); err == nil {
t.Errorf("packet for peer %d also delivered to peer %d", i, j)
}
}
}
}
// TestProxyDropsPacketsOutsideTheRange guards the wildcard bind: anything that
// is not addressed to a handed-out endpoint must not reach a relayed peer.
func TestProxyDropsPacketsOutsideTheRange(t *testing.T) {
proxy := NewProxy(testWGPort+1, 1280)
if err := proxy.Listen(); err != nil {
t.Fatalf("listen: %v", err)
}
defer func() {
if err := proxy.Free(); err != nil {
t.Errorf("free proxy: %v", err)
}
}()
proxySide, testSide := relayEnd(t)
if _, _, err := proxy.AddRelayedConn(proxySide); err != nil {
t.Fatalf("add relayed conn: %v", err)
}
sender, err := net.Dial("udp", net.JoinHostPort("127.0.0.1", strconv.Itoa(proxy.proxyPort)))
if err != nil {
t.Fatalf("sender: %v", err)
}
defer func() {
if err := sender.Close(); err != nil {
t.Logf("close sender: %v", err)
}
}()
if _, err := sender.Write([]byte("stray")); err != nil {
t.Fatalf("write stray packet: %v", err)
}
buf := make([]byte, 1500)
if err := testSide.SetReadDeadline(time.Now().Add(500 * time.Millisecond)); err != nil {
t.Fatalf("set read deadline: %v", err)
}
if _, _, err := testSide.ReadFrom(buf); err == nil {
t.Error("packet addressed to 127.0.0.1 was forwarded to a relayed peer")
}
}
// A wrapper that is closed before it starts forwarding still has to give its
// endpoint address back, otherwise the range leaks an address per attempt.
func TestClosingBeforeWorkReleasesTheAddress(t *testing.T) {
proxy := NewProxy(testWGPort+2, 1280)
if err := proxy.Listen(); err != nil {
t.Fatalf("listen: %v", err)
}
defer func() {
if err := proxy.Free(); err != nil {
t.Errorf("free proxy: %v", err)
}
}()
proxySide, _ := relayEnd(t)
wrapper := NewProxyWrapper(proxy)
if err := wrapper.AddRelayedConn(context.Background(), nil, proxySide); err != nil {
t.Fatalf("add relayed conn: %v", err)
}
if got := len(proxy.relayedConnStore); got != 1 {
t.Fatalf("store holds %d entries after adding one conn, want 1", got)
}
if err := wrapper.CloseConn(); err != nil {
t.Fatalf("close conn: %v", err)
}
if got := len(proxy.relayedConnStore); got != 0 {
t.Errorf("store holds %d entries after close, want 0", got)
}
}
@@ -1,6 +1,6 @@
//go:build linux && !android //go:build linux && !android
package ebpf package loopback
import ( import (
"context" "context"
@@ -8,6 +8,7 @@ import (
"fmt" "fmt"
"io" "io"
"net" "net"
"net/netip"
"sync" "sync"
"github.com/google/gopacket" "github.com/google/gopacket"
@@ -95,13 +96,14 @@ func NewPacketHeaders(localWGListenPort int, endpoint *net.UDPAddr) (*PacketHead
// ProxyWrapper help to keep the remoteConn instance for net.Conn.Close function call // ProxyWrapper help to keep the remoteConn instance for net.Conn.Close function call
type ProxyWrapper struct { type ProxyWrapper struct {
wgeBPFProxy *WGEBPFProxy proxy *Proxy
remoteConn net.Conn remoteConn net.Conn
ctx context.Context ctx context.Context
cancel context.CancelFunc cancel context.CancelFunc
wgRelayedEndpointAddr *net.UDPAddr wgRelayedEndpointAddr *net.UDPAddr
peerAddr netip.Addr
headers *PacketHeaders headers *PacketHeaders
headerCurrentUsed *PacketHeaders headerCurrentUsed *PacketHeaders
rawConn net.PacketConn rawConn net.PacketConn
@@ -113,36 +115,44 @@ type ProxyWrapper struct {
closeListener *listener.CloseListener closeListener *listener.CloseListener
} }
func NewProxyWrapper(proxy *WGEBPFProxy) *ProxyWrapper { func NewProxyWrapper(proxy *Proxy) *ProxyWrapper {
return &ProxyWrapper{ return &ProxyWrapper{
wgeBPFProxy: proxy, proxy: proxy,
pausedCond: sync.NewCond(&sync.Mutex{}), pausedCond: sync.NewCond(&sync.Mutex{}),
closeListener: listener.NewCloseListener(), closeListener: listener.NewCloseListener(),
} }
} }
func (p *ProxyWrapper) AddRelayedConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error { func (p *ProxyWrapper) AddRelayedConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error {
addr, err := p.wgeBPFProxy.AddRelayedConn(remoteConn) addr, peerAddr, err := p.proxy.AddRelayedConn(remoteConn)
if err != nil { if err != nil {
return fmt.Errorf("add relayed conn: %w", err) return fmt.Errorf("add relayed conn: %w", err)
} }
headers, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, addr) // the endpoint address is otherwise only released by the forwarding
// goroutine, which never starts when the setup below fails
release := func() { p.proxy.removeRelayedConn(peerAddr, remoteConn) }
headers, err := NewPacketHeaders(p.proxy.localWGListenPort, addr)
if err != nil { if err != nil {
release()
return fmt.Errorf("create packet sender: %w", err) return fmt.Errorf("create packet sender: %w", err)
} }
// Check if required raw connection is available // Check if required raw connection is available
if !headers.isIPv4 && p.wgeBPFProxy.rawConnIPv6 == nil { if !headers.isIPv4 && p.proxy.rawConnIPv6 == nil {
release()
return errIPv6ConnNotAvailable return errIPv6ConnNotAvailable
} }
if headers.isIPv4 && p.wgeBPFProxy.rawConnIPv4 == nil { if headers.isIPv4 && p.proxy.rawConnIPv4 == nil {
release()
return errIPv4ConnNotAvailable return errIPv4ConnNotAvailable
} }
p.remoteConn = remoteConn p.remoteConn = remoteConn
p.ctx, p.cancel = context.WithCancel(ctx) p.ctx, p.cancel = context.WithCancel(ctx)
p.wgRelayedEndpointAddr = addr p.wgRelayedEndpointAddr = addr
p.peerAddr = peerAddr
p.headers = headers p.headers = headers
p.rawConn = p.selectRawConn(headers) p.rawConn = p.selectRawConn(headers)
return nil return nil
@@ -193,18 +203,18 @@ func (p *ProxyWrapper) RedirectAs(endpoint *net.UDPAddr) {
return return
} }
header, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, endpoint) header, err := NewPacketHeaders(p.proxy.localWGListenPort, endpoint)
if err != nil { if err != nil {
log.Errorf("failed to create packet headers: %s", err) log.Errorf("failed to create packet headers: %s", err)
return return
} }
// Check if required raw connection is available // Check if required raw connection is available
if !header.isIPv4 && p.wgeBPFProxy.rawConnIPv6 == nil { if !header.isIPv4 && p.proxy.rawConnIPv6 == nil {
log.Error(errIPv6ConnNotAvailable) log.Error(errIPv6ConnNotAvailable)
return return
} }
if header.isIPv4 && p.wgeBPFProxy.rawConnIPv4 == nil { if header.isIPv4 && p.proxy.rawConnIPv4 == nil {
log.Error(errIPv4ConnNotAvailable) log.Error(errIPv4ConnNotAvailable)
return return
} }
@@ -240,6 +250,10 @@ func (p *ProxyWrapper) CloseConn() error {
p.closeListener.SetCloseListener(nil) p.closeListener.SetCloseListener(nil)
// releases the endpoint address for a wrapper that was never started, and
// is a no-op once the forwarding goroutine has released it
p.proxy.removeRelayedConn(p.peerAddr, p.remoteConn)
p.pausedCond.L.Lock() p.pausedCond.L.Lock()
p.paused = false p.paused = false
p.pausedCond.Signal() p.pausedCond.Signal()
@@ -252,9 +266,9 @@ func (p *ProxyWrapper) CloseConn() error {
} }
func (p *ProxyWrapper) proxyToLocal(ctx context.Context) { func (p *ProxyWrapper) proxyToLocal(ctx context.Context) {
defer p.wgeBPFProxy.removeRelayedConn(uint16(p.wgRelayedEndpointAddr.Port)) defer p.proxy.removeRelayedConn(p.peerAddr, p.remoteConn)
buf := make([]byte, p.wgeBPFProxy.mtu+bufsize.WGBufferOverhead) buf := make([]byte, p.proxy.mtu+bufsize.WGBufferOverhead)
for { for {
n, err := p.readFromRemote(ctx, buf) n, err := p.readFromRemote(ctx, buf)
if err != nil { if err != nil {
@@ -286,7 +300,7 @@ func (p *ProxyWrapper) readFromRemote(ctx context.Context, buf []byte) (int, err
} }
p.closeListener.Notify() p.closeListener.Notify()
if !errors.Is(err, io.EOF) { if !errors.Is(err, io.EOF) {
log.Errorf("failed to read from relayed conn (endpoint: :%d): %s", p.wgRelayedEndpointAddr.Port, err) log.Errorf("failed to read from relayed conn (endpoint: %s): %s", p.wgRelayedEndpointAddr, err)
} }
return 0, err return 0, err
} }
@@ -314,7 +328,7 @@ func (p *ProxyWrapper) sendPkg(data []byte, header *PacketHeaders) error {
func (p *ProxyWrapper) selectRawConn(header *PacketHeaders) net.PacketConn { func (p *ProxyWrapper) selectRawConn(header *PacketHeaders) net.PacketConn {
if header.isIPv4 { if header.isIPv4 {
return p.wgeBPFProxy.rawConnIPv4 return p.proxy.rawConnIPv4
} }
return p.wgeBPFProxy.rawConnIPv6 return p.proxy.rawConnIPv6
} }
+17 -17
View File
@@ -9,25 +9,25 @@ import (
"github.com/netbirdio/netbird/client/iface/bind" "github.com/netbirdio/netbird/client/iface/bind"
"github.com/netbirdio/netbird/client/iface/wgaddr" "github.com/netbirdio/netbird/client/iface/wgaddr"
bindproxy "github.com/netbirdio/netbird/client/iface/wgproxy/bind" bindproxy "github.com/netbirdio/netbird/client/iface/wgproxy/bind"
"github.com/netbirdio/netbird/client/iface/wgproxy/ebpf" "github.com/netbirdio/netbird/client/iface/wgproxy/loopback"
"github.com/netbirdio/netbird/client/iface/wgproxy/udp" "github.com/netbirdio/netbird/client/iface/wgproxy/udp"
) )
func seedProxies() ([]proxyInstance, error) { func seedProxies() ([]proxyInstance, error) {
pl := make([]proxyInstance, 0) pl := make([]proxyInstance, 0)
ebpfProxy := ebpf.NewWGEBPFProxy(51831, 1280) loopbackProxy := loopback.NewProxy(51831, 1280)
if err := ebpfProxy.Listen(); err != nil { if err := loopbackProxy.Listen(); err != nil {
return nil, fmt.Errorf("failed to initialize ebpf proxy: %s", err) return nil, fmt.Errorf("failed to initialize loopback proxy: %s", err)
} }
pEbpf := proxyInstance{ pLoopback := proxyInstance{
name: "ebpf kernel proxy", name: "loopback kernel proxy",
proxy: ebpf.NewProxyWrapper(ebpfProxy), proxy: loopback.NewProxyWrapper(loopbackProxy),
wgPort: 51831, wgPort: 51831,
closeFn: ebpfProxy.Free, closeFn: loopbackProxy.Free,
} }
pl = append(pl, pEbpf) pl = append(pl, pLoopback)
pUDP := proxyInstance{ pUDP := proxyInstance{
name: "udp kernel proxy", name: "udp kernel proxy",
@@ -42,18 +42,18 @@ func seedProxies() ([]proxyInstance, error) {
func seedProxyForProxyCloseByRemoteConn() ([]proxyInstance, error) { func seedProxyForProxyCloseByRemoteConn() ([]proxyInstance, error) {
pl := make([]proxyInstance, 0) pl := make([]proxyInstance, 0)
ebpfProxy := ebpf.NewWGEBPFProxy(51831, 1280) loopbackProxy := loopback.NewProxy(51831, 1280)
if err := ebpfProxy.Listen(); err != nil { if err := loopbackProxy.Listen(); err != nil {
return nil, fmt.Errorf("failed to initialize ebpf proxy: %s", err) return nil, fmt.Errorf("failed to initialize loopback proxy: %s", err)
} }
pEbpf := proxyInstance{ pLoopback := proxyInstance{
name: "ebpf kernel proxy", name: "loopback kernel proxy",
proxy: ebpf.NewProxyWrapper(ebpfProxy), proxy: loopback.NewProxyWrapper(loopbackProxy),
wgPort: 51831, wgPort: 51831,
closeFn: ebpfProxy.Free, closeFn: loopbackProxy.Free,
} }
pl = append(pl, pEbpf) pl = append(pl, pLoopback)
pUDP := proxyInstance{ pUDP := proxyInstance{
name: "udp kernel proxy", name: "udp kernel proxy",
+23 -23
View File
@@ -8,7 +8,7 @@ import (
"testing" "testing"
"time" "time"
"github.com/netbirdio/netbird/client/iface/wgproxy/ebpf" "github.com/netbirdio/netbird/client/iface/wgproxy/loopback"
"github.com/netbirdio/netbird/client/iface/wgproxy/udp" "github.com/netbirdio/netbird/client/iface/wgproxy/udp"
) )
@@ -198,20 +198,20 @@ func testRedirectAs(t *testing.T, proxy Proxy, wgPort int, nbAddr, p2pEndpoint *
} }
} }
// TestRedirectAs_eBPF_IPv4 tests RedirectAs with eBPF proxy using IPv4 addresses // TestRedirectAs_Loopback_IPv4 tests RedirectAs with the loopback proxy using IPv4 addresses
func TestRedirectAs_eBPF_IPv4(t *testing.T) { func TestRedirectAs_Loopback_IPv4(t *testing.T) {
wgPort := 51850 wgPort := 51850
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280) loopbackProxy := loopback.NewProxy(wgPort, 1280)
if err := ebpfProxy.Listen(); err != nil { if err := loopbackProxy.Listen(); err != nil {
t.Fatalf("failed to initialize ebpf proxy: %v", err) t.Fatalf("failed to initialize loopback proxy: %v", err)
} }
defer func() { defer func() {
if err := ebpfProxy.Free(); err != nil { if err := loopbackProxy.Free(); err != nil {
t.Errorf("failed to free ebpf proxy: %v", err) t.Errorf("failed to free loopback proxy: %v", err)
} }
}() }()
proxy := ebpf.NewProxyWrapper(ebpfProxy) proxy := loopback.NewProxyWrapper(loopbackProxy)
// NetBird UDP address of the remote peer // NetBird UDP address of the remote peer
nbAddr := &net.UDPAddr{ nbAddr := &net.UDPAddr{
@@ -227,20 +227,20 @@ func TestRedirectAs_eBPF_IPv4(t *testing.T) {
testRedirectAs(t, proxy, wgPort, nbAddr, p2pEndpoint) testRedirectAs(t, proxy, wgPort, nbAddr, p2pEndpoint)
} }
// TestRedirectAs_eBPF_IPv6 tests RedirectAs with eBPF proxy using IPv6 addresses // TestRedirectAs_Loopback_IPv6 tests RedirectAs with the loopback proxy using IPv6 addresses
func TestRedirectAs_eBPF_IPv6(t *testing.T) { func TestRedirectAs_Loopback_IPv6(t *testing.T) {
wgPort := 51851 wgPort := 51851
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280) loopbackProxy := loopback.NewProxy(wgPort, 1280)
if err := ebpfProxy.Listen(); err != nil { if err := loopbackProxy.Listen(); err != nil {
t.Fatalf("failed to initialize ebpf proxy: %v", err) t.Fatalf("failed to initialize loopback proxy: %v", err)
} }
defer func() { defer func() {
if err := ebpfProxy.Free(); err != nil { if err := loopbackProxy.Free(); err != nil {
t.Errorf("failed to free ebpf proxy: %v", err) t.Errorf("failed to free loopback proxy: %v", err)
} }
}() }()
proxy := ebpf.NewProxyWrapper(ebpfProxy) proxy := loopback.NewProxyWrapper(loopbackProxy)
// NetBird UDP address of the remote peer // NetBird UDP address of the remote peer
nbAddr := &net.UDPAddr{ nbAddr := &net.UDPAddr{
@@ -259,17 +259,17 @@ func TestRedirectAs_eBPF_IPv6(t *testing.T) {
// TestRedirectAs_Multiple_Switches tests switching between multiple endpoints // TestRedirectAs_Multiple_Switches tests switching between multiple endpoints
func TestRedirectAs_Multiple_Switches(t *testing.T) { func TestRedirectAs_Multiple_Switches(t *testing.T) {
wgPort := 51856 wgPort := 51856
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280) loopbackProxy := loopback.NewProxy(wgPort, 1280)
if err := ebpfProxy.Listen(); err != nil { if err := loopbackProxy.Listen(); err != nil {
t.Fatalf("failed to initialize ebpf proxy: %v", err) t.Fatalf("failed to initialize loopback proxy: %v", err)
} }
defer func() { defer func() {
if err := ebpfProxy.Free(); err != nil { if err := loopbackProxy.Free(); err != nil {
t.Errorf("failed to free ebpf proxy: %v", err) t.Errorf("failed to free loopback proxy: %v", err)
} }
}() }()
proxy := ebpf.NewProxyWrapper(ebpfProxy) proxy := loopback.NewProxyWrapper(loopbackProxy)
ctx := context.Background() ctx := context.Background()
+67 -3
View File
@@ -90,8 +90,9 @@ type StatusRecorder interface {
// fallback T-FinalWarningLead dialog (suppressed when the user dismissed // fallback T-FinalWarningLead dialog (suppressed when the user dismissed
// the first one for the same deadline). Safe for concurrent use. // the first one for the same deadline). Safe for concurrent use.
type Watcher struct { type Watcher struct {
lead time.Duration lead time.Duration
finalLead time.Duration finalLead time.Duration
deadlineOnly bool
mu sync.Mutex mu sync.Mutex
current time.Time current time.Time
@@ -102,6 +103,7 @@ type Watcher struct {
dismissedAt time.Time // deadline value the user dismissed via Dismiss(); gates fireFinal dismissedAt time.Time // deadline value the user dismissed via Dismiss(); gates fireFinal
closed bool closed bool
recorder StatusRecorder recorder StatusRecorder
nowFn func() time.Time
} }
// New returns a watcher with the package defaults WarningLead and // New returns a watcher with the package defaults WarningLead and
@@ -122,9 +124,17 @@ func NewWithLeads(lead, final time.Duration, recorder StatusRecorder) *Watcher {
lead: lead, lead: lead,
finalLead: final, finalLead: final,
recorder: recorder, recorder: recorder,
nowFn: time.Now,
} }
} }
// NewDeadlineOnly returns a watcher that validates and records deadlines but arms no warning timers.
func NewDeadlineOnly(recorder StatusRecorder) *Watcher {
w := New(recorder)
w.deadlineOnly = true
return w
}
// Update sets the latest deadline. Pass the zero time to clear (e.g. when // Update sets the latest deadline. Pass the zero time to clear (e.g. when
// a Sync push from the server omits the field because login expiration // a Sync push from the server omits the field because login expiration
// was disabled). // was disabled).
@@ -181,7 +191,7 @@ func (w *Watcher) Update(deadline time.Time) error {
w.finalFiredAt = time.Time{} w.finalFiredAt = time.Time{}
w.dismissedAt = time.Time{} w.dismissedAt = time.Time{}
if deadline.After(now) { if deadline.After(now) && !w.deadlineOnly {
w.armTimerLocked(deadline) w.armTimerLocked(deadline)
} }
recorder := w.recorder recorder := w.recorder
@@ -303,6 +313,11 @@ func (w *Watcher) fire(armedFor time.Time) {
w.mu.Unlock() w.mu.Unlock()
return return
} }
now := w.nowFn()
if isLate(now, armedFor, max(w.finalLead, 0)) {
w.fireLateLocked(armedFor, now)
return
}
w.firedAt = armedFor w.firedAt = armedFor
recorder := w.recorder recorder := w.recorder
w.mu.Unlock() w.mu.Unlock()
@@ -331,6 +346,14 @@ func (w *Watcher) fireFinal(armedFor time.Time) {
log.Infof("auth session final-warning skipped (dismissed by user)") log.Infof("auth session final-warning skipped (dismissed by user)")
return return
} }
now := w.nowFn()
if isLate(now, armedFor, 0) {
w.finalFiredAt = armedFor
w.mu.Unlock()
log.Infof("auth session final-warning skipped for deadline %s (passed %s ago)",
armedFor.Format(time.RFC3339), now.Round(0).Sub(armedFor).Round(time.Second))
return
}
w.finalFiredAt = armedFor w.finalFiredAt = armedFor
recorder := w.recorder recorder := w.recorder
w.mu.Unlock() w.mu.Unlock()
@@ -341,6 +364,39 @@ func (w *Watcher) fireFinal(armedFor time.Time) {
publishWarning(recorder, armedFor, true) publishWarning(recorder, armedFor, true)
} }
// fireLateLocked handles a T-WarningLead callback that fired inside the
// final-warning window: it sends the final warning in its place while the
// deadline has not passed and the user has not dismissed it, so a resume
// with time left still warns. The caller must hold w.mu; this helper
// releases it.
func (w *Watcher) fireLateLocked(armedFor, now time.Time) {
w.firedAt = armedFor
switch {
case w.dismissedAt.Equal(armedFor):
w.mu.Unlock()
log.Infof("auth session expiry soon warning skipped (dismissed by user)")
return
case w.finalFiredAt.Equal(armedFor):
w.mu.Unlock()
log.Infof("auth session expiry soon warning skipped (final warning already fired)")
return
case isLate(now, armedFor, 0):
w.mu.Unlock()
log.Infof("auth session expiry soon warning skipped for deadline %s (passed %s ago)",
armedFor.Format(time.RFC3339), now.Round(0).Sub(armedFor).Round(time.Second))
return
}
w.finalFiredAt = armedFor
recorder := w.recorder
w.mu.Unlock()
if recorder == nil {
return
}
log.Infof("auth session expiry soon warning fired inside the final-warning window, sending final warning for deadline %s",
armedFor.Format(time.RFC3339))
publishWarning(recorder, armedFor, true)
}
// armOneShotLocked schedules cb at fireAt. When fireAt is already in the // armOneShotLocked schedules cb at fireAt. When fireAt is already in the
// past it dispatches on the next scheduler tick so a state-change recorder // past it dispatches on the next scheduler tick so a state-change recorder
// notification (invoked after w.mu is released) lands first. Caller must // notification (invoked after w.mu is released) lands first. Caller must
@@ -380,3 +436,11 @@ func publishWarning(recorder StatusRecorder, deadline time.Time, final bool) {
meta, meta,
) )
} }
// isLate reports whether the wall clock now has already reached armedFor
// minus cutoffLead. The timers run on the monotonic clock, which can stall
// while the host sleeps, so a timer can fire long after the window it was
// armed for.
func isLate(now, armedFor time.Time, cutoffLead time.Duration) bool {
return !now.Round(0).Before(armedFor.Add(-cutoffLead).Round(0))
}
@@ -527,3 +527,201 @@ func TestDismissBeforeUpdateIsNoop(t *testing.T) {
} }
t.Fatalf("final-warning did not publish after no-op pre-Update Dismiss, events=%+v", r.snapshot()) t.Fatalf("final-warning did not publish after no-op pre-Update Dismiss, events=%+v", r.snapshot())
} }
func TestIsLate(t *testing.T) {
armedFor := time.Date(2026, 10, 1, 12, 0, 0, 0, time.UTC)
lead := 2 * time.Minute
tests := []struct {
name string
now time.Time
cutoffLead time.Duration
want bool
}{
{"before cutoff", armedFor.Add(-3 * time.Minute), lead, false},
{"at cutoff", armedFor.Add(-lead), lead, true},
{"after cutoff", armedFor.Add(-time.Minute), lead, true},
{"zero lead before deadline", armedFor.Add(-time.Second), 0, false},
{"zero lead at deadline", armedFor, 0, true},
{"zero lead after deadline", armedFor.Add(time.Second), 0, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := isLate(tt.now, armedFor, tt.cutoffLead); got != tt.want {
t.Fatalf("isLate(%s, %s, %s) = %v, want %v", tt.now, armedFor, tt.cutoffLead, got, tt.want)
}
})
}
}
func TestIsLateIgnoresMonotonicReading(t *testing.T) {
now := time.Now()
wallOnly := now.Round(0)
if isLate(now, wallOnly.Add(time.Second), 0) {
t.Fatalf("now with monotonic reading must compare as wall clock before a later wall-only deadline")
}
if !isLate(now, wallOnly, 0) {
t.Fatalf("now with monotonic reading must compare as wall clock at an equal wall-only deadline")
}
}
func TestLateTimerFiring(t *testing.T) {
tests := []struct {
name string
final bool
beforeDl time.Duration
wantWarns int
wantFinals int
}{
{"warning on resume inside window", false, 3 * time.Minute, 1, 0},
{"warning promoted to final inside final window", false, time.Minute, 0, 1},
{"warning skipped past deadline", false, -time.Minute, 0, 0},
{"final on resume before deadline", true, time.Minute, 0, 1},
{"final skipped past deadline", true, -time.Minute, 0, 0},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
// The deadline is an hour out so the real timers never fire
// during the test; the late callback is invoked directly with an
// injected clock that simulates a resume near the deadline.
d := time.Now().Add(time.Hour).Round(0)
w.nowFn = func() time.Time { return d.Add(-tt.beforeDl) }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
if tt.final {
w.fireFinal(d)
} else {
w.fire(d)
}
events := r.snapshot()
if got := countWhere(events, event.isWarning); got != tt.wantWarns {
t.Fatalf("expected %d warning publishes, got %d: %+v", tt.wantWarns, got, events)
}
if got := countWhere(events, event.isFinalWarning); got != tt.wantFinals {
t.Fatalf("expected %d final-warning publishes, got %d: %+v", tt.wantFinals, got, events)
}
})
}
}
func TestPromotedFinalWarningIsNotRepeated(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
d := time.Now().Add(time.Hour).Round(0)
now := d.Add(-time.Minute)
w.nowFn = func() time.Time { return now }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
w.fire(d)
// The final timer was suspended too, so it fires even later than the
// warning timer, here still just before the deadline.
now = d.Add(-30 * time.Second)
w.fireFinal(d)
events := r.snapshot()
if got := countWhere(events, event.isFinalWarning); got != 1 {
t.Fatalf("expected exactly 1 final-warning publish, got %d: %+v", got, events)
}
if got := countWhere(events, event.isWarning); got != 0 {
t.Fatalf("expected no regular warning publish, got %d: %+v", got, events)
}
}
func TestPromotionRespectsDismiss(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
d := time.Now().Add(time.Hour).Round(0)
w.nowFn = func() time.Time { return d.Add(-time.Minute) }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
w.Dismiss()
w.fire(d)
events := r.snapshot()
if got := countWhere(events, func(e event) bool { return e.kind == publish }); got != 0 {
t.Fatalf("expected no publish after dismiss, got %d: %+v", got, events)
}
}
func TestPromotionSkippedWhenFinalAlreadyFired(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
// Both timers fall in the past after a long suspend and are dispatched
// with a zero delay, so the final callback can run before the warning one.
d := time.Now().Add(time.Hour).Round(0)
w.nowFn = func() time.Time { return d.Add(-time.Minute) }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
w.fireFinal(d)
w.fire(d)
events := r.snapshot()
if got := countWhere(events, event.isFinalWarning); got != 1 {
t.Fatalf("expected exactly 1 final-warning publish, got %d: %+v", got, events)
}
if got := countWhere(events, event.isWarning); got != 0 {
t.Fatalf("expected no regular warning publish, got %d: %+v", got, events)
}
}
func TestDeadlineOnlyRecordsDeadlineWithoutWarnings(t *testing.T) {
r := &fakeRecorder{}
w := NewDeadlineOnly(r)
defer w.Close()
// With the default leads this deadline would otherwise fire both
// timers on the next tick.
d := time.Now().Add(50 * time.Millisecond).Round(0)
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
if got := r.deadline(); !got.Equal(d) {
t.Fatalf("expected recorder deadline %v, got %v", d, got)
}
time.Sleep(100 * time.Millisecond)
events := r.snapshot()
if got := countWhere(events, func(e event) bool { return e.kind == publish }); got != 0 {
t.Fatalf("expected no publish in deadline-only mode, got %d: %+v", got, events)
}
if w.timer != nil || w.finalTimer != nil {
t.Fatal("expected no timers armed in deadline-only mode")
}
}
func TestDeadlineOnlyStillRejectsOutOfRangeDeadlines(t *testing.T) {
r := &fakeRecorder{}
w := NewDeadlineOnly(r)
defer w.Close()
if err := w.Update(time.Now().Add(time.Hour)); err != nil {
t.Fatalf("Update: %v", err)
}
err := w.Update(time.Now().Add(-maxPastHorizon - time.Hour))
if !errors.Is(err, ErrDeadlineInPast) {
t.Fatalf("expected ErrDeadlineInPast, got %v", err)
}
if got := r.deadline(); !got.IsZero() {
t.Fatalf("expected recorder cleared after rejection, got %v", got)
}
}
+42
View File
@@ -0,0 +1,42 @@
package daemonaddr
import (
"os"
"strconv"
log "github.com/sirupsen/logrus"
)
const (
// EnvMaxRecvMsgSize overrides the default gRPC max receive message size for
// connections to the daemon. Value is in bytes.
EnvMaxRecvMsgSize = "NB_DAEMON_GRPC_MAX_MSG_SIZE"
// defaultMaxRecvMsgSize is the max gRPC receive message size used for daemon
// connections when EnvMaxRecvMsgSize is unset or invalid. It overrides the
// gRPC library default of 4 MB, which a detailed status already exceeds on a
// network of a few thousand peers.
defaultMaxRecvMsgSize = 1024 * 1024 * 16
)
// MaxRecvMsgSize returns the max gRPC receive message size for daemon connections
// from the environment, or defaultMaxRecvMsgSize (16 MB) if unset or invalid.
func MaxRecvMsgSize() int {
val := os.Getenv(EnvMaxRecvMsgSize)
if val == "" {
return defaultMaxRecvMsgSize
}
size, err := strconv.Atoi(val)
if err != nil {
log.Warnf("invalid %s value %q, using default: %v", EnvMaxRecvMsgSize, val, err)
return defaultMaxRecvMsgSize
}
if size <= 0 {
log.Warnf("invalid %s value %d, must be positive, using default", EnvMaxRecvMsgSize, size)
return defaultMaxRecvMsgSize
}
return size
}
+112
View File
@@ -0,0 +1,112 @@
package daemonaddr
import (
"context"
"net"
"os"
"strings"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/credentials/insecure"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/proto"
)
func TestMaxRecvMsgSize(t *testing.T) {
tests := []struct {
name string
envValue string
expected int
}{
{name: "unset returns default", envValue: "", expected: defaultMaxRecvMsgSize},
{name: "non-numeric returns default", envValue: "abc", expected: defaultMaxRecvMsgSize},
{name: "negative returns default", envValue: "-1", expected: defaultMaxRecvMsgSize},
{name: "zero returns default", envValue: "0", expected: defaultMaxRecvMsgSize},
{name: "valid value is used", envValue: "33554432", expected: 33554432},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
// Set first so the previous value is restored on cleanup, then unset to
// exercise the absent case.
t.Setenv(EnvMaxRecvMsgSize, tc.envValue)
if tc.envValue == "" {
require.NoError(t, os.Unsetenv(EnvMaxRecvMsgSize), "unset the override")
}
assert.Equal(t, tc.expected, MaxRecvMsgSize(), "max receive message size")
})
}
}
// bigStatusServer answers Status with a response larger than gRPC's 4 MB default
// receive limit, which is what a detailed status on a large network looks like.
type bigStatusServer struct {
proto.UnimplementedDaemonServiceServer
payload string
}
func (s *bigStatusServer) Status(context.Context, *proto.StatusRequest) (*proto.StatusResponse, error) {
return &proto.StatusResponse{Status: s.payload}, nil
}
func startBigStatusServer(t *testing.T, payload string) string {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err, "listen on loopback")
srv := grpc.NewServer()
proto.RegisterDaemonServiceServer(srv, &bigStatusServer{payload: payload})
go func() {
_ = srv.Serve(listener)
}()
t.Cleanup(srv.Stop)
return "tcp://" + listener.Addr().String()
}
func TestDialTargetAcceptsAStatusOverTheGrpcDefault(t *testing.T) {
payload := strings.Repeat("x", 5*1024*1024)
addr := startBigStatusServer(t, payload)
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
target, opts := DialTarget(addr)
conn, err := grpc.NewClient(target, opts...)
require.NoError(t, err, "dial the daemon")
t.Cleanup(func() { _ = conn.Close() })
resp, err := proto.NewDaemonServiceClient(conn).Status(ctx, &proto.StatusRequest{})
require.NoError(t, err, "a detailed status must not be rejected for its size")
assert.Len(t, resp.GetStatus(), len(payload), "the whole response must arrive")
}
// TestDialTargetRaisesTheDefaultLimit is the negative control: the same response
// over a connection carrying gRPC's own defaults is refused, which is the failure
// reported by `netbird status -d` on a large deployment.
func TestDialTargetRaisesTheDefaultLimit(t *testing.T) {
payload := strings.Repeat("x", 5*1024*1024)
addr := startBigStatusServer(t, payload)
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
conn, err := grpc.NewClient(
strings.TrimPrefix(addr, "tcp://"),
grpc.WithTransportCredentials(insecure.NewCredentials()),
)
require.NoError(t, err, "dial with the library defaults")
t.Cleanup(func() { _ = conn.Close() })
_, err = proto.NewDaemonServiceClient(conn).Status(ctx, &proto.StatusRequest{})
require.Error(t, err, "the library default must reject this response")
assert.Equal(t, codes.ResourceExhausted, status.Code(err), "gRPC rejects an oversized message")
}
+4 -1
View File
@@ -36,7 +36,10 @@ const (
// address. The npipe scheme needs a context dialer because gRPC has no // address. The npipe scheme needs a context dialer because gRPC has no
// named-pipe resolver; unix and tcp are handled by gRPC itself. // named-pipe resolver; unix and tcp are handled by gRPC itself.
func DialTarget(addr string) (string, []grpc.DialOption) { func DialTarget(addr string) (string, []grpc.DialOption) {
opts := []grpc.DialOption{grpc.WithTransportCredentials(insecure.NewCredentials())} opts := []grpc.DialOption{
grpc.WithTransportCredentials(insecure.NewCredentials()),
grpc.WithDefaultCallOptions(grpc.MaxCallRecvMsgSize(MaxRecvMsgSize())),
}
if name, ok := strings.CutPrefix(addr, pipeScheme); ok { if name, ok := strings.CutPrefix(addr, pipeScheme); ok {
paths := PipePaths(name) paths := PipePaths(name)
+53 -1
View File
@@ -379,9 +379,38 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen
} }
} }
// bundleFilePattern names the bundle zips Generate creates in tempDir; the
// asterisk is filled in by os.CreateTemp.
const bundleFilePattern = "netbird.debug.*.zip"
const exportedBundlePrefix = "netbird.debug-file."
const exportedBundleMaxAge = 24 * time.Hour
// RemoveStaleBundles deletes bundle zips that an interrupted generation or
// upload left behind in dir. Only files older than maxAge go, so a bundle that
// another caller is still writing or uploading in the same directory survives.
// Exported bundles are kept for exportedBundleMaxAge instead.
func RemoveStaleBundles(dir string, maxAge time.Duration) {
removeStaleFiles(dir, bundleFilePattern, maxAge)
removeStaleFiles(dir, exportedBundlePrefix+"*.zip", exportedBundleMaxAge)
}
// ExportBundle renames a generated bundle out of the RemoveStaleBundles pattern
// and returns the new path. The caller owns the file from then on; an export
// abandoned for longer than exportedBundleMaxAge is removed by RemoveStaleBundles.
func ExportBundle(path string) (string, error) {
base := strings.TrimPrefix(filepath.Base(path), strings.SplitN(bundleFilePattern, "*", 2)[0])
exported := filepath.Join(filepath.Dir(path), exportedBundlePrefix+base)
if err := os.Rename(path, exported); err != nil {
return "", fmt.Errorf("export debug bundle: %w", err)
}
return exported, nil
}
// Generate creates a debug bundle and returns the location. // Generate creates a debug bundle and returns the location.
func (g *BundleGenerator) Generate() (resp string, err error) { func (g *BundleGenerator) Generate() (resp string, err error) {
bundlePath, err := os.CreateTemp(g.tempDir, "netbird.debug.*.zip") bundlePath, err := os.CreateTemp(g.tempDir, bundleFilePattern)
if err != nil { if err != nil {
return "", fmt.Errorf("create zip file: %w", err) return "", fmt.Errorf("create zip file: %w", err)
} }
@@ -1725,3 +1754,26 @@ func anonymizeSlice(v []any, anonymizer *anonymize.Anonymizer) []any {
} }
return v return v
} }
func removeStaleFiles(dir, pattern string, maxAge time.Duration) {
matches, err := filepath.Glob(filepath.Join(dir, pattern))
if err != nil {
log.Debugf("glob stale debug bundles in %s: %v", dir, err)
return
}
cutoff := time.Now().Add(-maxAge)
for _, path := range matches {
info, err := os.Stat(path)
if err != nil || info.ModTime().After(cutoff) {
continue
}
if err := os.Remove(path); err != nil {
if !errors.Is(err, fs.ErrNotExist) {
log.Warnf("remove stale debug bundle %s: %v", path, err)
}
continue
}
log.Infof("removed stale debug bundle %s", path)
}
}
+50
View File
@@ -4,6 +4,7 @@ import (
"archive/zip" "archive/zip"
"bytes" "bytes"
"encoding/json" "encoding/json"
"fmt"
"net" "net"
"net/netip" "net/netip"
"net/url" "net/url"
@@ -969,3 +970,52 @@ func renderAddConfigSpecific(g *BundleGenerator) string {
func newAnonymizerForTest() *anonymize.Anonymizer { func newAnonymizerForTest() *anonymize.Anonymizer {
return anonymize.NewAnonymizer(anonymize.DefaultAddresses()) return anonymize.NewAnonymizer(anonymize.DefaultAddresses())
} }
func TestRemoveStaleBundles(t *testing.T) {
dir := t.TempDir()
stale := filepath.Join(dir, "netbird.debug.111.zip")
fresh := filepath.Join(dir, "netbird.debug.222.zip")
other := filepath.Join(dir, "netbird.debug.333.txt")
owned := filepath.Join(dir, "netbird.debug.444.zip")
abandoned := filepath.Join(dir, "netbird.debug.555.zip")
for _, p := range []string{stale, fresh, other, owned, abandoned} {
require.NoError(t, os.WriteFile(p, []byte("x"), 0o600))
}
exported, err := ExportBundle(owned)
require.NoError(t, err)
exportedAbandoned, err := ExportBundle(abandoned)
require.NoError(t, err)
old := time.Now().Add(-2 * time.Hour)
for _, p := range []string{stale, other, exported} {
require.NoError(t, os.Chtimes(p, old, old))
}
ancient := time.Now().Add(-exportedBundleMaxAge - time.Hour)
require.NoError(t, os.Chtimes(exportedAbandoned, ancient, ancient))
RemoveStaleBundles(dir, time.Hour)
assert.NoFileExists(t, stale, "bundle older than maxAge should be removed")
assert.FileExists(t, fresh, "bundle younger than maxAge must survive, it may still be uploading")
assert.FileExists(t, other, "files outside the bundle pattern must not be touched")
assert.NoFileExists(t, owned)
assert.FileExists(t, exported, "exported bundle is caller-owned and must survive maxAge")
assert.NoFileExists(t, exportedAbandoned, "exported bundle older than exportedBundleMaxAge is abandoned")
}
func TestBundleIncludesNetworkMap(t *testing.T) {
for _, anonymize := range []bool{false, true} {
t.Run(fmt.Sprintf("anonymize=%t", anonymize), func(t *testing.T) {
g := NewBundleGenerator(GeneratorDependencies{
SyncResponse: &mgmProto.SyncResponse{NetworkMap: &mgmProto.NetworkMap{Serial: 1}},
}, BundleConfig{Anonymize: anonymize})
require.Contains(t, bundleEntries(t, g), "network_map.json")
})
}
}
func TestBundleOmitsNetworkMapWithoutSyncResponse(t *testing.T) {
g := NewBundleGenerator(GeneratorDependencies{}, BundleConfig{})
require.NotContains(t, bundleEntries(t, g), "network_map.json")
}
+84 -14
View File
@@ -124,19 +124,9 @@ func newHostManager(wgInterface WGIface) (*registryConfigurator, error) {
return nil, err return nil, err
} }
var useGPO bool
k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE)
if err != nil {
log.Debugf("failed to open GPO DNS policy root: %v", err)
} else {
closer(k)
useGPO = true
log.Infof("detected GPO DNS policy configuration, using policy store")
}
configurator := &registryConfigurator{ configurator := &registryConfigurator{
guid: guid, guid: guid,
gpo: useGPO, gpo: useGPOPolicyStore(),
} }
origNameservers, err := configurator.captureOriginalNameservers() origNameservers, err := configurator.captureOriginalNameservers()
@@ -576,14 +566,22 @@ func (r *registryConfigurator) setInterfaceRegistryKeyStringValue(key, value str
return nil return nil
} }
// deleteInterfaceRegistryKeyProperty removes a value from the interface key.
// A value that is already gone, or an interface key that is, is not an error:
// the caller asked for the value not to be there, and a cleanup that runs twice
// has to reach its later steps on the second run as well.
func (r *registryConfigurator) deleteInterfaceRegistryKeyProperty(propertyKey string) error { func (r *registryConfigurator) deleteInterfaceRegistryKeyProperty(propertyKey string) error {
regKey, err := r.getInterfaceRegistryKey() regKey, err := r.getInterfaceRegistryKey()
if err != nil { switch {
case errors.Is(err, registry.ErrNotExist), errors.Is(err, syscall.ERROR_PATH_NOT_FOUND):
log.Debugf("interface key of %s does not exist, nothing to delete %s from", r.guid, propertyKey)
return nil
case err != nil:
return fmt.Errorf("get interface registry key: %w", err) return fmt.Errorf("get interface registry key: %w", err)
} }
defer closer(regKey) defer closer(regKey)
if err := regKey.DeleteValue(propertyKey); err != nil { if err := regKey.DeleteValue(propertyKey); err != nil && !errors.Is(err, registry.ErrNotExist) {
return fmt.Errorf("delete registry key %s: %w", propertyKey, err) return fmt.Errorf("delete registry key %s: %w", propertyKey, err)
} }
return nil return nil
@@ -612,7 +610,12 @@ func (r *registryConfigurator) restoreHostDNS() error {
go r.flushDNSCache() go r.flushDNSCache()
return nil // Last, and only on the way out, once no rule of ours is left: during a
// session the store is where the rules of this run live, and emptying it
// mid-session would have the next rule recreate it anyway. Propagated so a
// failure keeps the shutdown state for the next run to retry, rather than
// leaving the store to hold up every rule change from here on.
return removeEmptyGPOPolicyStore()
} }
// removeDNSMatchPolicies deletes every NRPT rule this client may have created, // removeDNSMatchPolicies deletes every NRPT rule this client may have created,
@@ -651,6 +654,73 @@ func (r *registryConfigurator) restoreUncleanShutdownDNS() error {
return r.restoreHostDNS() return r.restoreHostDNS()
} }
// useGPOPolicyStore reports whether NRPT rules have to go into the group policy
// store, and clears an empty one out of the way first.
//
// The order is the point. A store left empty by an earlier run would otherwise
// decide this run too, sending its rules somewhere the resolver only reads when
// the policy engine next applies DNS client policy. Removing it before the
// choice is made leaves the local store authoritative for the whole session,
// including the first one after an upgrade.
func useGPOPolicyStore() bool {
if err := removeEmptyGPOPolicyStore(); err != nil {
// Nothing to retry against here: the worst case is the run going
// through the group policy store, which is where it would have gone
// before this check existed.
log.Warnf("%v", err)
}
k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE)
if err != nil {
log.Debugf("failed to open GPO DNS policy root: %v", err)
return false
}
closer(k)
log.Infof("detected GPO DNS policy configuration, using policy store")
return true
}
// removeEmptyGPOPolicyStore deletes the group policy DnsPolicyConfig key once
// nothing is left in it. The key survives the deletion of the last rule it
// held, and the client treats its presence as "group policy configures the
// NRPT", so an empty one left behind keeps every later run writing rules there.
// Rules in that store reach the resolver only when the policy engine next
// applies DNS client policy, and a rule this client writes belongs to no GPO,
// so nothing schedules that application: both adding and removing a rule are
// held up by a minute or more, and for a removal that is a catch-all rule
// resolving every name over an interface that no longer exists. With the store
// absent the local one is authoritative and a change applies at once.
//
// A store that still holds rules, values or subkeys of somebody else's is left
// alone.
func removeEmptyGPOPolicyStore() error {
k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE)
switch {
case errors.Is(err, registry.ErrNotExist), errors.Is(err, syscall.ERROR_PATH_NOT_FOUND):
return nil
case err != nil:
return fmt.Errorf("open HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err)
}
info, err := k.Stat()
closer(k)
if err != nil {
return fmt.Errorf("stat HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err)
}
if info.SubKeyCount != 0 || info.ValueCount != 0 {
return nil
}
if err := registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot); err != nil {
return fmt.Errorf("delete empty HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err)
}
log.Infof("removed the empty GPO DNS policy store, leaving the local one authoritative")
return nil
}
// listNRPTRuleKeys returns the names of our NRPT rule keys under a policy store // listNRPTRuleKeys returns the names of our NRPT rule keys under a policy store
// root. An absent root holds nothing to clean up, which is the normal state of // root. An absent root holds nothing to clean up, which is the normal state of
// the GPO store on a machine without DNS Client policy. // the GPO store on a machine without DNS Client policy.
+129
View File
@@ -8,6 +8,8 @@ import (
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"golang.org/x/sys/windows/registry" "golang.org/x/sys/windows/registry"
"github.com/netbirdio/netbird/client/internal/winregistry"
) )
// TestNRPTEntriesCleanupOnConfigChange tests that old NRPT entries are properly cleaned up // TestNRPTEntriesCleanupOnConfigChange tests that old NRPT entries are properly cleaned up
@@ -405,3 +407,130 @@ func TestNRPTDomainBatching(t *testing.T) {
}) })
} }
} }
// TestRemoveEmptyGPOPolicyStore verifies that cleanup takes the GPO policy
// store itself with it once our rules are gone, since the store existing keeps
// the local one from being applied, and that a store with somebody else's rule
// in it is left alone.
func TestRemoveEmptyGPOPolicyStore(t *testing.T) {
if testing.Short() {
t.Skip("skipping registry integration test in short mode")
}
t.Cleanup(func() { cleanupRegistryKeys(t) })
cleanupRegistryKeys(t)
testIP := netip.MustParseAddr("100.64.0.1")
cfg := &registryConfigurator{gpo: true}
// a store holding a rule of ours is kept, because the rule is still applied
require.NoError(t, cfg.addDNSMatchPolicy([]string{".example.com"}, testIP))
exists, err := registryKeyExists(gpoDnsPolicyConfigMatchPath + "-0")
require.NoError(t, err)
require.True(t, exists, "Should write the rule to the GPO policy store")
require.NoError(t, removeEmptyGPOPolicyStore())
exists, err = registryKeyExists(GPODNSPolicyConfigRoot)
require.NoError(t, err)
assert.True(t, exists, "Should keep a policy store that still holds a rule")
// once the rules are gone the store goes with them
require.NoError(t, cfg.removeDNSMatchPolicies())
require.NoError(t, removeEmptyGPOPolicyStore())
exists, err = registryKeyExists(GPODNSPolicyConfigRoot)
require.NoError(t, err)
assert.False(t, exists, "Should remove the GPO policy store once it is empty")
// A store is not ours to remove while somebody else has a rule in it. The
// rule is written volatile like our own: the rules above created the parent
// chain volatile, and Windows refuses a stable subkey under a volatile
// parent.
foreignRule := GPODNSPolicyConfigRoot + `\{2A3B4C5D-6E7F-4041-8283-84858687888A}`
foreignKey, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, foreignRule, registry.SET_VALUE)
require.NoError(t, err, "Should create a foreign GPO rule")
foreignKey.Close()
t.Cleanup(func() {
_ = registry.DeleteKey(registry.LOCAL_MACHINE, foreignRule)
_ = registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot)
})
require.NoError(t, cfg.removeDNSMatchPolicies())
require.NoError(t, removeEmptyGPOPolicyStore())
exists, err = registryKeyExists(foreignRule)
require.NoError(t, err)
assert.True(t, exists, "Should not remove a foreign rule")
exists, err = registryKeyExists(GPODNSPolicyConfigRoot)
require.NoError(t, err)
assert.True(t, exists, "Should keep a policy store that still holds a foreign rule")
}
// TestDeleteInterfaceRegistryKeyPropertyTwice verifies that removing a value
// that is already gone, or one on an interface key that is, reports success.
// Teardown runs again after a failed cleanup, and the steps that follow this
// one have to be reached on that second run.
func TestDeleteInterfaceRegistryKeyPropertyTwice(t *testing.T) {
if testing.Short() {
t.Skip("skipping registry integration test in short mode")
}
testGUID := "{12345678-1234-1234-1234-123456789ABC}"
interfacePath := InterfaceConfigPath + `\` + testGUID
testKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, interfacePath, registry.SET_VALUE)
require.NoError(t, err, "Should create test interface registry key")
testKey.Close()
t.Cleanup(func() {
_ = registry.DeleteKey(registry.LOCAL_MACHINE, interfacePath)
})
cfg := &registryConfigurator{guid: testGUID}
require.NoError(t, cfg.setInterfaceRegistryKeyStringValue(interfaceConfigSearchListKey, "example.com"))
require.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey))
assert.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey),
"Should report success for a value that is already gone")
// and with the interface key itself gone, as it is once the adapter is
require.NoError(t, registry.DeleteKey(registry.LOCAL_MACHINE, interfacePath))
assert.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey),
"Should report success when the interface key does not exist")
}
// TestUseGPOPolicyStoreClearsEmptyStore verifies that the store is cleared
// before it is consulted, so an empty one left by an earlier run does not send
// this run's rules to the group policy store. A store somebody else has a rule
// in still decides where the rules go.
func TestUseGPOPolicyStoreClearsEmptyStore(t *testing.T) {
if testing.Short() {
t.Skip("skipping registry integration test in short mode")
}
t.Cleanup(func() { cleanupRegistryKeys(t) })
cleanupRegistryKeys(t)
// the leftover an earlier run used to keep, which the client read as
// "group policy configures the NRPT" for every run after it
emptyStore, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.SET_VALUE)
require.NoError(t, err, "Should create the GPO policy store")
emptyStore.Close()
assert.False(t, useGPOPolicyStore(), "An empty store should not decide where the rules go")
exists, err := registryKeyExists(GPODNSPolicyConfigRoot)
require.NoError(t, err)
assert.False(t, exists, "Should clear the empty store before consulting it")
foreignRule := GPODNSPolicyConfigRoot + `\{2A3B4C5D-6E7F-4041-8283-84858687888A}`
foreignKey, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, foreignRule, registry.SET_VALUE)
require.NoError(t, err, "Should create a foreign GPO rule")
foreignKey.Close()
t.Cleanup(func() {
_ = registry.DeleteKey(registry.LOCAL_MACHINE, foreignRule)
_ = registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot)
})
assert.True(t, useGPOPolicyStore(), "A store holding a rule should decide where the rules go")
exists, err = registryKeyExists(GPODNSPolicyConfigRoot)
require.NoError(t, err)
assert.True(t, exists, "Should keep a store that holds a rule")
}
+7 -10
View File
@@ -9,9 +9,9 @@ import (
"os" "os"
"testing" "testing"
"go.uber.org/mock/gomock"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"go.uber.org/mock/gomock"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes" "golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"github.com/netbirdio/netbird/client/iface" "github.com/netbirdio/netbird/client/iface"
@@ -24,6 +24,10 @@ import (
nbdns "github.com/netbirdio/netbird/dns" nbdns "github.com/netbirdio/netbird/dns"
) )
// testIFaceBlackList mirrors the overlay prefixes profilemanager.DefaultInterfaceBlacklist
// carries. Declared here rather than imported because profilemanager imports this package.
var testIFaceBlackList = []string{"wt", "utun", "tun0"}
func TestUpdateDNSServer(t *testing.T) { func TestUpdateDNSServer(t *testing.T) {
nameServers := []nbdns.NameServer{ nameServers := []nbdns.NameServer{
@@ -243,10 +247,7 @@ func TestUpdateDNSServer(t *testing.T) {
for n, testCase := range testCases { for n, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) { t.Run(testCase.name, func(t *testing.T) {
privKey, _ := wgtypes.GenerateKey() privKey, _ := wgtypes.GenerateKey()
newNet, err := stdnet.NewNet(context.Background(), nil) newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
if err != nil {
t.Fatal(err)
}
opts := iface.WGIFaceOpts{ opts := iface.WGIFaceOpts{
IFaceName: fmt.Sprintf("utun230%d", n), IFaceName: fmt.Sprintf("utun230%d", n),
@@ -348,11 +349,7 @@ func TestDNSFakeResolverHandleUpdates(t *testing.T) {
defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) defer t.Setenv("NB_WG_KERNEL_DISABLED", ov)
t.Setenv("NB_WG_KERNEL_DISABLED", "true") t.Setenv("NB_WG_KERNEL_DISABLED", "true")
newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"}) newNet := stdnet.NewNet(context.Background(), []string{"utun2301"})
if err != nil {
t.Errorf("create stdnet: %v", err)
return
}
privKey, _ := wgtypes.GeneratePrivateKey() privKey, _ := wgtypes.GeneratePrivateKey()
opts := iface.WGIFaceOpts{ opts := iface.WGIFaceOpts{
+1 -5
View File
@@ -394,11 +394,7 @@ func createWgInterfaceWithBind(t *testing.T) (*iface.WGIface, error) {
defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) defer t.Setenv("NB_WG_KERNEL_DISABLED", ov)
t.Setenv("NB_WG_KERNEL_DISABLED", "true") t.Setenv("NB_WG_KERNEL_DISABLED", "true")
newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"}) newNet := stdnet.NewNet(context.Background(), []string{"utun2301"})
if err != nil {
t.Fatalf("create stdnet: %v", err)
return nil, err
}
privKey, _ := wgtypes.GeneratePrivateKey() privKey, _ := wgtypes.GeneratePrivateKey()
-148
View File
@@ -1,148 +0,0 @@
// Code generated by bpf2go; DO NOT EDIT.
//go:build mips || mips64 || ppc64 || s390x
package ebpf
import (
"bytes"
_ "embed"
"fmt"
"io"
"github.com/cilium/ebpf"
)
// loadBpf returns the embedded CollectionSpec for bpf.
func loadBpf() (*ebpf.CollectionSpec, error) {
reader := bytes.NewReader(_BpfBytes)
spec, err := ebpf.LoadCollectionSpecFromReader(reader)
if err != nil {
return nil, fmt.Errorf("can't load bpf: %w", err)
}
return spec, err
}
// loadBpfObjects loads bpf and converts it into a struct.
//
// The following types are suitable as obj argument:
//
// *bpfObjects
// *bpfPrograms
// *bpfMaps
//
// See ebpf.CollectionSpec.LoadAndAssign documentation for details.
func loadBpfObjects(obj interface{}, opts *ebpf.CollectionOptions) error {
spec, err := loadBpf()
if err != nil {
return err
}
return spec.LoadAndAssign(obj, opts)
}
// bpfSpecs contains maps and programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfSpecs struct {
bpfProgramSpecs
bpfMapSpecs
bpfVariableSpecs
}
// bpfProgramSpecs contains programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfProgramSpecs struct {
NbXdpProg *ebpf.ProgramSpec `ebpf:"nb_xdp_prog"`
}
// bpfMapSpecs contains maps before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfMapSpecs struct {
NbFeatures *ebpf.MapSpec `ebpf:"nb_features"`
NbWgProxySettingsMap *ebpf.MapSpec `ebpf:"nb_wg_proxy_settings_map"`
}
// bpfVariableSpecs contains global variables before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfVariableSpecs struct {
FlagFeatureWgProxy *ebpf.VariableSpec `ebpf:"flag_feature_wg_proxy"`
MapKeyFeatures *ebpf.VariableSpec `ebpf:"map_key_features"`
MapKeyProxyPort *ebpf.VariableSpec `ebpf:"map_key_proxy_port"`
MapKeyWgPort *ebpf.VariableSpec `ebpf:"map_key_wg_port"`
ProxyPort *ebpf.VariableSpec `ebpf:"proxy_port"`
WgPort *ebpf.VariableSpec `ebpf:"wg_port"`
}
// bpfObjects contains all objects after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfObjects struct {
bpfPrograms
bpfMaps
bpfVariables
}
func (o *bpfObjects) Close() error {
return _BpfClose(
&o.bpfPrograms,
&o.bpfMaps,
)
}
// bpfMaps contains all maps after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfMaps struct {
NbFeatures *ebpf.Map `ebpf:"nb_features"`
NbWgProxySettingsMap *ebpf.Map `ebpf:"nb_wg_proxy_settings_map"`
}
func (m *bpfMaps) Close() error {
return _BpfClose(
m.NbFeatures,
m.NbWgProxySettingsMap,
)
}
// bpfVariables contains all global variables after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfVariables struct {
FlagFeatureWgProxy *ebpf.Variable `ebpf:"flag_feature_wg_proxy"`
MapKeyFeatures *ebpf.Variable `ebpf:"map_key_features"`
MapKeyProxyPort *ebpf.Variable `ebpf:"map_key_proxy_port"`
MapKeyWgPort *ebpf.Variable `ebpf:"map_key_wg_port"`
ProxyPort *ebpf.Variable `ebpf:"proxy_port"`
WgPort *ebpf.Variable `ebpf:"wg_port"`
}
// bpfPrograms contains all programs after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfPrograms struct {
NbXdpProg *ebpf.Program `ebpf:"nb_xdp_prog"`
}
func (p *bpfPrograms) Close() error {
return _BpfClose(
p.NbXdpProg,
)
}
func _BpfClose(closers ...io.Closer) error {
for _, closer := range closers {
if err := closer.Close(); err != nil {
return err
}
}
return nil
}
// Do not access this directly.
//
//go:embed bpf_bpfeb.o
var _BpfBytes []byte
Binary file not shown.
-148
View File
@@ -1,148 +0,0 @@
// Code generated by bpf2go; DO NOT EDIT.
//go:build 386 || amd64 || arm || arm64 || loong64 || mips64le || mipsle || ppc64le || riscv64 || wasm
package ebpf
import (
"bytes"
_ "embed"
"fmt"
"io"
"github.com/cilium/ebpf"
)
// loadBpf returns the embedded CollectionSpec for bpf.
func loadBpf() (*ebpf.CollectionSpec, error) {
reader := bytes.NewReader(_BpfBytes)
spec, err := ebpf.LoadCollectionSpecFromReader(reader)
if err != nil {
return nil, fmt.Errorf("can't load bpf: %w", err)
}
return spec, err
}
// loadBpfObjects loads bpf and converts it into a struct.
//
// The following types are suitable as obj argument:
//
// *bpfObjects
// *bpfPrograms
// *bpfMaps
//
// See ebpf.CollectionSpec.LoadAndAssign documentation for details.
func loadBpfObjects(obj interface{}, opts *ebpf.CollectionOptions) error {
spec, err := loadBpf()
if err != nil {
return err
}
return spec.LoadAndAssign(obj, opts)
}
// bpfSpecs contains maps and programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfSpecs struct {
bpfProgramSpecs
bpfMapSpecs
bpfVariableSpecs
}
// bpfProgramSpecs contains programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfProgramSpecs struct {
NbXdpProg *ebpf.ProgramSpec `ebpf:"nb_xdp_prog"`
}
// bpfMapSpecs contains maps before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfMapSpecs struct {
NbFeatures *ebpf.MapSpec `ebpf:"nb_features"`
NbWgProxySettingsMap *ebpf.MapSpec `ebpf:"nb_wg_proxy_settings_map"`
}
// bpfVariableSpecs contains global variables before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfVariableSpecs struct {
FlagFeatureWgProxy *ebpf.VariableSpec `ebpf:"flag_feature_wg_proxy"`
MapKeyFeatures *ebpf.VariableSpec `ebpf:"map_key_features"`
MapKeyProxyPort *ebpf.VariableSpec `ebpf:"map_key_proxy_port"`
MapKeyWgPort *ebpf.VariableSpec `ebpf:"map_key_wg_port"`
ProxyPort *ebpf.VariableSpec `ebpf:"proxy_port"`
WgPort *ebpf.VariableSpec `ebpf:"wg_port"`
}
// bpfObjects contains all objects after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfObjects struct {
bpfPrograms
bpfMaps
bpfVariables
}
func (o *bpfObjects) Close() error {
return _BpfClose(
&o.bpfPrograms,
&o.bpfMaps,
)
}
// bpfMaps contains all maps after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfMaps struct {
NbFeatures *ebpf.Map `ebpf:"nb_features"`
NbWgProxySettingsMap *ebpf.Map `ebpf:"nb_wg_proxy_settings_map"`
}
func (m *bpfMaps) Close() error {
return _BpfClose(
m.NbFeatures,
m.NbWgProxySettingsMap,
)
}
// bpfVariables contains all global variables after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfVariables struct {
FlagFeatureWgProxy *ebpf.Variable `ebpf:"flag_feature_wg_proxy"`
MapKeyFeatures *ebpf.Variable `ebpf:"map_key_features"`
MapKeyProxyPort *ebpf.Variable `ebpf:"map_key_proxy_port"`
MapKeyWgPort *ebpf.Variable `ebpf:"map_key_wg_port"`
ProxyPort *ebpf.Variable `ebpf:"proxy_port"`
WgPort *ebpf.Variable `ebpf:"wg_port"`
}
// bpfPrograms contains all programs after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfPrograms struct {
NbXdpProg *ebpf.Program `ebpf:"nb_xdp_prog"`
}
func (p *bpfPrograms) Close() error {
return _BpfClose(
p.NbXdpProg,
)
}
func _BpfClose(closers ...io.Closer) error {
for _, closer := range closers {
if err := closer.Close(); err != nil {
return err
}
}
return nil
}
// Do not access this directly.
//
//go:embed bpf_bpfel.o
var _BpfBytes []byte
Binary file not shown.
-115
View File
@@ -1,115 +0,0 @@
package ebpf
import (
_ "embed"
"net"
"sync"
"github.com/cilium/ebpf/link"
"github.com/cilium/ebpf/rlimit"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/ebpf/manager"
)
const (
mapKeyFeatures uint32 = 0
featureFlagWGProxy = 0b00000001
)
var (
singleton manager.Manager
singletonLock = &sync.Mutex{}
)
// required packages libbpf-dev, libc6-dev-i386-amd64-cross
// GeneralManager is used to load multiple eBPF programs with a custom check (if then) done in prog.c
// The manager simply adds a feature (byte) of each program to a map that is shared between the userspace and kernel.
// When packet arrives, the C code checks for each feature (if it is set) and executes each enabled program (e.g., wg_proxy.c).
//
//go:generate go run github.com/cilium/ebpf/cmd/bpf2go -cc clang-14 bpf src/prog.c -- -I /usr/x86_64-linux-gnu/include -include src/bpf_map_def.h
type GeneralManager struct {
lock sync.Mutex
link link.Link
featureFlags uint16
bpfObjs bpfObjects
}
// GetEbpfManagerInstance return a static eBpf Manager instance
func GetEbpfManagerInstance() manager.Manager {
singletonLock.Lock()
defer singletonLock.Unlock()
if singleton != nil {
return singleton
}
singleton = &GeneralManager{}
return singleton
}
func (tf *GeneralManager) setFeatureFlag(feature uint16) {
tf.featureFlags |= feature
}
func (tf *GeneralManager) loadXdp() error {
if tf.link != nil {
return nil
}
// it required for Docker
err := rlimit.RemoveMemlock()
if err != nil {
return err
}
iFace, err := net.InterfaceByName("lo")
if err != nil {
return err
}
// load pre-compiled programs into the kernel.
err = loadBpfObjects(&tf.bpfObjs, nil)
if err != nil {
return err
}
tf.link, err = link.AttachXDP(link.XDPOptions{
Program: tf.bpfObjs.NbXdpProg,
Interface: iFace.Index,
})
if err != nil {
_ = tf.bpfObjs.Close()
tf.link = nil
return err
}
return nil
}
func (tf *GeneralManager) unsetFeatureFlag(feature uint16) error {
tf.lock.Lock()
defer tf.lock.Unlock()
tf.featureFlags &^= feature
if tf.link == nil {
return nil
}
if tf.featureFlags == 0 {
return tf.close()
}
return tf.bpfObjs.NbFeatures.Put(mapKeyFeatures, tf.featureFlags)
}
func (tf *GeneralManager) close() error {
log.Debugf("detach ebpf program ")
err := tf.bpfObjs.Close()
if err != nil {
log.Warnf("failed to close eBpf objects: %s", err)
}
err = tf.link.Close()
tf.link = nil
return err
}
@@ -1,31 +0,0 @@
package ebpf
import (
"testing"
)
func TestManager_setFeatureFlag(t *testing.T) {
mgr := GeneralManager{}
mgr.setFeatureFlag(featureFlagWGProxy)
if mgr.featureFlags != featureFlagWGProxy {
t.Errorf("invalid feature state")
}
mgr.setFeatureFlag(featureFlagWGProxy)
if mgr.featureFlags != featureFlagWGProxy {
t.Errorf("setting a flag twice must be idempotent, got: %d", mgr.featureFlags)
}
}
func TestManager_unsetFeatureFlag(t *testing.T) {
mgr := GeneralManager{}
mgr.setFeatureFlag(featureFlagWGProxy)
err := mgr.unsetFeatureFlag(featureFlagWGProxy)
if err != nil {
t.Errorf("unexpected error: %s", err)
}
if mgr.featureFlags != 0 {
t.Errorf("invalid feature state, expected: %d, got: %d", 0, mgr.featureFlags)
}
}
@@ -1,16 +0,0 @@
// libbpf 1.0 removed struct bpf_map_def, but the programs here keep the legacy
// map definitions: they load on kernels built without BTF, which BTF-style
// (SEC(".maps")) definitions do not. Define the struct ourselves so the
// programs compile against current libbpf headers.
#ifndef NB_BPF_MAP_DEF_H
#define NB_BPF_MAP_DEF_H
struct bpf_map_def {
unsigned int type;
unsigned int key_size;
unsigned int value_size;
unsigned int max_entries;
unsigned int map_flags;
};
#endif
-54
View File
@@ -1,54 +0,0 @@
#include <stdbool.h>
#include <linux/if_ether.h> // ETH_P_IP
#include <linux/udp.h>
#include <linux/ip.h>
#include <netinet/in.h>
#include <linux/bpf.h>
#include <bpf/bpf_helpers.h>
#include "wg_proxy.c"
const __u16 flag_feature_wg_proxy = 0b01;
const __u32 map_key_features = 0;
struct bpf_map_def SEC("maps") nb_features = {
.type = BPF_MAP_TYPE_ARRAY,
.key_size = sizeof(__u32),
.value_size = sizeof(__u16),
.max_entries = 10,
};
SEC("xdp")
int nb_xdp_prog(struct xdp_md *ctx) {
__u16 *features;
features = bpf_map_lookup_elem(&nb_features, &map_key_features);
if (!features) {
return XDP_PASS;
}
void *data = (void *)(long)ctx->data;
void *data_end = (void *)(long)ctx->data_end;
struct ethhdr *eth = data;
struct iphdr *ip = (data + sizeof(struct ethhdr));
struct udphdr *udp = (data + sizeof(struct ethhdr) + sizeof(struct iphdr));
// return early if not enough data
if (data + sizeof(struct ethhdr) + sizeof(struct iphdr) + sizeof(struct udphdr) > data_end){
return XDP_PASS;
}
// skip non IPv4 packages
if (eth->h_proto != htons(ETH_P_IP)) {
return XDP_PASS;
}
// skip non UPD packages
if (ip->protocol != IPPROTO_UDP) {
return XDP_PASS;
}
if (*features & flag_feature_wg_proxy) {
xdp_wg_proxy(ip, udp);
}
return XDP_PASS;
}
char _license[] SEC("license") = "GPL";
-27
View File
@@ -1,27 +0,0 @@
# XDP programs
`prog.c` is attached to the `lo` device and dispatches to the features enabled in the
`nb_features` map. The only feature is the WireGuard proxy (`wg_proxy.c`): it rewrites
loopback UDP sent from the WireGuard listen port so it reaches the userspace relay proxy
port instead, and swaps the peer endpoint port into the source so the proxy can tell
peers apart.
Maps use the legacy `struct bpf_map_def` form, defined in `bpf_map_def.h` because libbpf
1.0 removed it. They load on kernels built without BTF, which BTF-style (`SEC(".maps")`)
definitions do not.
Regenerate the objects with `go generate ./client/internal/ebpf/ebpf/`; it needs
`clang-14`. Loading a regenerated object needs root, attaching it needs `bpf_link`
(kernel >= 5.7), and only one XDP program can own `lo` at a time.
# Debug
The CONFIG_BPF_EVENTS kernel module is required for bpf_printk.
Apply this code to use bpf_printk
```
#define bpf_printk(fmt, ...) \
({ \
char ____fmt[] = fmt; \
bpf_trace_printk(____fmt, sizeof(____fmt), ##__VA_ARGS__); \
})
```
-60
View File
@@ -1,60 +0,0 @@
const __u32 map_key_proxy_port = 0;
const __u32 map_key_wg_port = 1;
struct bpf_map_def SEC("maps") nb_wg_proxy_settings_map = {
.type = BPF_MAP_TYPE_ARRAY,
.key_size = sizeof(__u32),
.value_size = sizeof(__u16),
.max_entries = 10,
};
__u16 proxy_port = 0;
__u16 wg_port = 0;
bool read_port_settings() {
__u16 *value;
value = bpf_map_lookup_elem(&nb_wg_proxy_settings_map, &map_key_proxy_port);
if (!value) {
return false;
}
proxy_port = *value;
value = bpf_map_lookup_elem(&nb_wg_proxy_settings_map, &map_key_wg_port);
if (!value) {
return false;
}
wg_port = htons(*value);
return true;
}
int xdp_wg_proxy(struct iphdr *ip, struct udphdr *udp) {
if (proxy_port == 0 || wg_port == 0) {
if (!read_port_settings()){
return XDP_PASS;
}
// bpf_printk("proxy port: %d, wg port: %d", proxy_port, wg_port);
}
// 2130706433 = 127.0.0.1
if (ip->daddr != htonl(2130706433)) {
return XDP_PASS;
}
if (udp->source != wg_port){
return XDP_PASS;
}
__be16 new_src_port = udp->dest;
__be16 new_dst_port = htons(proxy_port);
udp->dest = new_dst_port;
udp->source = new_src_port;
// The ports are covered by the UDP checksum. This is an IPv4 loopback hop
// and the payload is already integrity-protected, so clear the checksum (a
// zero UDP checksum means "not computed" for IPv4) rather than leave a
// stale value the kernel would drop as UDP_CSUM.
udp->check = 0;
return XDP_PASS;
}
@@ -1,41 +0,0 @@
package ebpf
import log "github.com/sirupsen/logrus"
const (
mapKeyProxyPort uint32 = 0
mapKeyWgPort uint32 = 1
)
func (tf *GeneralManager) LoadWgProxy(proxyPort, wgPort int) error {
log.Debugf("load ebpf WG proxy")
tf.lock.Lock()
defer tf.lock.Unlock()
err := tf.loadXdp()
if err != nil {
return err
}
err = tf.bpfObjs.NbWgProxySettingsMap.Put(mapKeyProxyPort, uint16(proxyPort))
if err != nil {
return err
}
err = tf.bpfObjs.NbWgProxySettingsMap.Put(mapKeyWgPort, uint16(wgPort))
if err != nil {
return err
}
tf.setFeatureFlag(featureFlagWGProxy)
err = tf.bpfObjs.NbFeatures.Put(mapKeyFeatures, tf.featureFlags)
if err != nil {
return err
}
return nil
}
func (tf *GeneralManager) FreeWGProxy() error {
log.Debugf("free ebpf WG proxy")
return tf.unsetFeatureFlag(featureFlagWGProxy)
}
@@ -1,15 +0,0 @@
//go:build !android
package ebpf
import (
"github.com/netbirdio/netbird/client/internal/ebpf/ebpf"
"github.com/netbirdio/netbird/client/internal/ebpf/manager"
)
// GetEbpfManagerInstance is a wrapper function. This encapsulation is required because if the code import the internal
// ebpf package the Go compiler will include the object files. But it is not supported on Android. It can cause instant
// panic on older Android version.
func GetEbpfManagerInstance() manager.Manager {
return ebpf.GetEbpfManagerInstance()
}
@@ -1,10 +0,0 @@
//go:build !linux || android
package ebpf
import "github.com/netbirdio/netbird/client/internal/ebpf/manager"
// GetEbpfManagerInstance return error because ebpf is not supported on all os
func GetEbpfManagerInstance() manager.Manager {
panic("unsupported os")
}
-7
View File
@@ -1,7 +0,0 @@
package manager
// Manager is used to load multiple eBPF programs. E.g., the WireGuard proxy
type Manager interface {
LoadWgProxy(proxyPort, wgPort int) error
FreeWGProxy() error
}
+11
View File
@@ -6,6 +6,17 @@ import (
"path/filepath" "path/filepath"
) )
// CheckOnlyOwnerWritable reports an error unless path, and every directory
// leading to it, is owned by an account that can already act with the privileges
// the caller holds, and is writable by nobody else.
//
// Exported for callers outside elevation that read a file while privileged and
// then act on what it says: the same question this package asks of an
// executable, asked of a configuration file.
func CheckOnlyOwnerWritable(path string) error {
return checkOnlyOwnerWritable(path)
}
// trustedSelf returns the path of this executable, provided it is one we are // trustedSelf returns the path of this executable, provided it is one we are
// willing to have run as root. // willing to have run as root.
// //
+6 -26
View File
@@ -664,10 +664,6 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL)
} }
e.wgDevice.Store(e.wgInterface.GetWGDevice()) e.wgDevice.Store(e.wgInterface.GetWGDevice())
// Set up notrack rules immediately after proxy is listening to prevent
// conntrack entries from being created before the rules are in place
e.setupWGProxyNoTrack()
// Start after interface is up since port may have been resolved from 0 or changed if occupied // Start after interface is up since port may have been resolved from 0 or changed if occupied
e.shutdownWg.Add(1) e.shutdownWg.Add(1)
go func() { go func() {
@@ -805,23 +801,6 @@ func (e *Engine) initFirewall() error {
return nil return nil
} }
// setupWGProxyNoTrack configures connection tracking exclusion for WireGuard proxy traffic.
// This prevents conntrack/MASQUERADE from affecting loopback traffic between WireGuard and the eBPF proxy.
func (e *Engine) setupWGProxyNoTrack() {
if e.firewall == nil {
return
}
proxyPort := e.wgInterface.GetProxyPort()
if proxyPort == 0 {
return
}
if err := e.firewall.SetupEBPFProxyNoTrack(proxyPort, uint16(e.config.WgPort)); err != nil {
log.Warnf("failed to setup ebpf proxy notrack: %v", err)
}
}
func (e *Engine) blockLanAccess() { func (e *Engine) blockLanAccess() {
if e.config.BlockInbound { if e.config.BlockInbound {
// no need to set up extra deny rules if inbound is already blocked in general // no need to set up extra deny rules if inbound is already blocked in general
@@ -1064,7 +1043,11 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
// back to empty if the FQDN doesn't have the expected shape. // back to empty if the FQDN doesn't have the expected shape.
dnsName = extractDNSDomainFromFQDN(pc.GetFqdn()) dnsName = extractDNSDomainFromFQDN(pc.GetFqdn())
} }
result, err := nbnetworkmap.EnvelopeToNetworkMap(e.ctx, envelope, localKey, dnsName) // With the firewall disabled there is no ACL manager to program, so
// RoutesFirewallRules would be built and then dropped. On a peer that
// routes many network resources that is the single most expensive
// step of the sync.
result, err := nbnetworkmap.EnvelopeToNetworkMap(e.ctx, envelope, localKey, dnsName, e.config.DisableFirewall)
if err != nil { if err != nil {
return fmt.Errorf("decode network map envelope: %w", err) return fmt.Errorf("decode network map envelope: %w", err)
} }
@@ -2208,10 +2191,7 @@ func (e *Engine) close() {
} }
func (e *Engine) newWgIface() (*iface.WGIface, error) { func (e *Engine) newWgIface() (*iface.WGIface, error) {
transportNet, err := e.newStdNet() transportNet := e.newStdNet()
if err != nil {
log.Errorf("failed to create pion's stdnet: %s", err)
}
opts := iface.WGIFaceOpts{ opts := iface.WGIFaceOpts{
IFaceName: e.config.WgIfaceName, IFaceName: e.config.WgIfaceName,
+15 -11
View File
@@ -12,12 +12,12 @@ import (
"testing" "testing"
"time" "time"
"go.uber.org/mock/gomock"
"github.com/google/uuid" "github.com/google/uuid"
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"go.opentelemetry.io/otel" "go.opentelemetry.io/otel"
"go.uber.org/mock/gomock"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes" "golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"google.golang.org/grpc" "google.golang.org/grpc"
"google.golang.org/grpc/keepalive" "google.golang.org/grpc/keepalive"
@@ -27,6 +27,7 @@ import (
"github.com/netbirdio/netbird/client/iface/wgaddr" "github.com/netbirdio/netbird/client/iface/wgaddr"
"github.com/netbirdio/netbird/client/internal/dns" "github.com/netbirdio/netbird/client/internal/dns"
"github.com/netbirdio/netbird/client/internal/peer" "github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/profilemanager"
nbssh "github.com/netbirdio/netbird/client/ssh" nbssh "github.com/netbirdio/netbird/client/ssh"
"github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/client/system"
nbdns "github.com/netbirdio/netbird/dns" nbdns "github.com/netbirdio/netbird/dns"
@@ -81,6 +82,7 @@ func TestEngine_SSH(t *testing.T) {
WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"), WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"),
WgPrivateKey: key, WgPrivateKey: key,
WgPort: 33100, WgPort: 33100,
IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
ServerSSHAllowed: true, ServerSSHAllowed: true,
MTU: iface.DefaultMTU, MTU: iface.DefaultMTU,
SSHKey: sshKey, SSHKey: sshKey,
@@ -204,11 +206,12 @@ func TestEngine_Sync(t *testing.T) {
} }
relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU) relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU)
engine := NewEngine(ctx, cancel, &EngineConfig{ engine := NewEngine(ctx, cancel, &EngineConfig{
WgIfaceName: "utun103", WgIfaceName: "utun103",
WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"), WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"),
WgPrivateKey: key, WgPrivateKey: key,
WgPort: 33100, WgPort: 33100,
MTU: iface.DefaultMTU, IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
MTU: iface.DefaultMTU,
}, EngineServices{ }, EngineServices{
SignalClient: &signal.MockClient{}, SignalClient: &signal.MockClient{},
MgmClient: &mgmt.MockClient{SyncFunc: syncFunc}, MgmClient: &mgmt.MockClient{SyncFunc: syncFunc},
@@ -412,11 +415,12 @@ func createEngine(ctx context.Context, cancel context.CancelFunc, setupKey strin
wgPort := 33100 + i wgPort := 33100 + i
conf := &EngineConfig{ conf := &EngineConfig{
WgIfaceName: ifaceName, WgIfaceName: ifaceName,
WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address), WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address),
WgPrivateKey: key, WgPrivateKey: key,
WgPort: wgPort, WgPort: wgPort,
MTU: iface.DefaultMTU, IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
MTU: iface.DefaultMTU,
} }
relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU) relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU)
+7 -5
View File
@@ -1,4 +1,4 @@
//go:build !js //go:build !js && !android
package internal package internal
@@ -7,10 +7,12 @@ import (
"github.com/netbirdio/netbird/client/internal/peer" "github.com/netbirdio/netbird/client/internal/peer"
) )
// newSessionWatcher returns the real SSO session expiry watcher for every // newSessionWatcher returns the real SSO session expiry watcher. The js/wasm
// non-wasm build. The js/wasm build gets a no-op stub from // build gets a no-op stub from engine_sessionwatch_js.go so the sessionwatch
// engine_sessionwatch_js.go so the sessionwatch package (and its timer // package (and its timer machinery) never links into the wasm binary; the
// machinery) never links into the wasm binary. // android build gets a deadline-only watcher from
// engine_sessionwatch_android.go because the app schedules the warnings
// itself.
func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher { func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher {
return sessionwatch.New(recorder) return sessionwatch.New(recorder)
} }
@@ -0,0 +1,12 @@
//go:build android
package internal
import (
"github.com/netbirdio/netbird/client/internal/auth/sessionwatch"
"github.com/netbirdio/netbird/client/internal/peer"
)
func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher {
return sessionwatch.NewDeadlineOnly(recorder)
}
+1 -1
View File
@@ -6,6 +6,6 @@ import (
"github.com/netbirdio/netbird/client/internal/stdnet" "github.com/netbirdio/netbird/client/internal/stdnet"
) )
func (e *Engine) newStdNet() (*stdnet.Net, error) { func (e *Engine) newStdNet() *stdnet.Net {
return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList) return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList)
} }
+1 -1
View File
@@ -2,6 +2,6 @@ package internal
import "github.com/netbirdio/netbird/client/internal/stdnet" import "github.com/netbirdio/netbird/client/internal/stdnet"
func (e *Engine) newStdNet() (*stdnet.Net, error) { func (e *Engine) newStdNet() *stdnet.Net {
return stdnet.NewNetWithDiscover(e.clientCtx, e.mobileDep.IFaceDiscover, e.config.IFaceBlackList) return stdnet.NewNetWithDiscover(e.clientCtx, e.mobileDep.IFaceDiscover, e.config.IFaceBlackList)
} }
+20 -16
View File
@@ -65,7 +65,6 @@ type MockWGIface struct {
GetStatsFunc func() (map[string]configurer.WGStats, error) GetStatsFunc func() (map[string]configurer.WGStats, error)
GetInterfaceGUIDStringFunc func() (string, error) GetInterfaceGUIDStringFunc func() (string, error)
GetProxyFunc func() wgproxy.Proxy GetProxyFunc func() wgproxy.Proxy
GetProxyPortFunc func() uint16
GetNetFunc func() *netstack.Net GetNetFunc func() *netstack.Net
LastActivitiesFunc func() map[string]monotime.Time LastActivitiesFunc func() map[string]monotime.Time
} }
@@ -162,13 +161,6 @@ func (m *MockWGIface) GetProxy() wgproxy.Proxy {
return m.GetProxyFunc() return m.GetProxyFunc()
} }
func (m *MockWGIface) GetProxyPort() uint16 {
if m.GetProxyPortFunc != nil {
return m.GetProxyPortFunc()
}
return 0
}
func (m *MockWGIface) GetNet() *netstack.Net { func (m *MockWGIface) GetNet() *netstack.Net {
return m.GetNetFunc() return m.GetNetFunc()
} }
@@ -696,10 +688,7 @@ func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) {
StatusRecorder: peer.NewRecorder("https://mgm"), StatusRecorder: peer.NewRecorder("https://mgm"),
}, MobileDependency{}) }, MobileDependency{})
engine.ctx = ctx engine.ctx = ctx
newNet, err := stdnet.NewNet(context.Background(), nil) newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
if err != nil {
t.Fatal(err)
}
opts := iface.WGIFaceOpts{ opts := iface.WGIFaceOpts{
IFaceName: wgIfaceName, IFaceName: wgIfaceName,
@@ -904,10 +893,7 @@ func TestEngine_UpdateNetworkMapWithDNSUpdate(t *testing.T) {
}, MobileDependency{}) }, MobileDependency{})
engine.ctx = ctx engine.ctx = ctx
newNet, err := stdnet.NewNet(context.Background(), nil) newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
if err != nil {
t.Fatal(err)
}
opts := iface.WGIFaceOpts{ opts := iface.WGIFaceOpts{
IFaceName: wgIfaceName, IFaceName: wgIfaceName,
Address: wgaddr.MustParseWGAddress(wgAddr), Address: wgaddr.MustParseWGAddress(wgAddr),
@@ -1533,3 +1519,21 @@ func TestOverlayAddrsFromAllowedIPs(t *testing.T) {
}) })
} }
} }
func TestEngine_SyncResponsePersistence(t *testing.T) {
e := &Engine{}
_, err := e.GetLatestSyncResponse()
require.Error(t, err, "persistence is disabled by default")
e.SetSyncResponsePersistence(true)
e.persistSyncResponse(&mgmtProto.SyncResponse{NetworkMap: &mgmtProto.NetworkMap{Serial: 7}})
got, err := e.GetLatestSyncResponse()
require.NoError(t, err)
assert.Equal(t, uint64(7), got.GetNetworkMap().GetSerial())
e.SetSyncResponsePersistence(false)
_, err = e.GetLatestSyncResponse()
require.Error(t, err)
}
-1
View File
@@ -28,7 +28,6 @@ type wgIfaceBase interface {
Up() (*udpmux.UniversalUDPMuxDefault, error) Up() (*udpmux.UniversalUDPMuxDefault, error)
UpdateAddr(newAddr wgaddr.Address) error UpdateAddr(newAddr wgaddr.Address) error
GetProxy() wgproxy.Proxy GetProxy() wgproxy.Proxy
GetProxyPort() uint16
UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error
RemoveEndpointAddress(key string) error RemoveEndpointAddress(key string) error
RemovePeer(peerKey string) error RemovePeer(peerKey string) error
+1 -1
View File
@@ -38,7 +38,7 @@ func asDaemon(t *testing.T, id Identity) {
prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate
t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate }) t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate })
selfIdentity, selfKnown = id, true selfIdentity, selfKnown = id, true
selfMayDelegate = !id.IsPrivileged() selfMayDelegate = mayDelegate(id)
} }
func TestCallerIdentity_DirectConnections(t *testing.T) { func TestCallerIdentity_DirectConnections(t *testing.T) {
+6 -6
View File
@@ -18,7 +18,8 @@ import (
"google.golang.org/grpc/peer" "google.golang.org/grpc/peer"
) )
// Well-known Windows SIDs that identify a fully privileged principal. // Well-known Windows SIDs. Only LocalSystem and BUILTIN\Administrators identify a
// privileged principal; the service accounts are shared by unrelated services.
const ( const (
sidLocalSystem = "S-1-5-18" // NT AUTHORITY\SYSTEM sidLocalSystem = "S-1-5-18" // NT AUTHORITY\SYSTEM
sidLocalService = "S-1-5-19" // NT AUTHORITY\LOCAL SERVICE sidLocalService = "S-1-5-19" // NT AUTHORITY\LOCAL SERVICE
@@ -67,9 +68,9 @@ func (i Identity) IsWindows() bool {
// user-to-root boundary. // user-to-root boundary.
// //
// On Windows the decision comes from the caller's token rather than from // On Windows the decision comes from the caller's token rather than from
// account names or group RIDs: an elevated token, one of the service accounts // account names or group RIDs: an elevated token, the LocalSystem SID, or a
// the daemon itself may run as, or a token with BUILTIN\Administrators // token with BUILTIN\Administrators enabled. LocalService and NetworkService
// enabled. A UAC-filtered administrator has that group marked deny-only, and // are not privileged by SID. A UAC-filtered administrator has that group marked deny-only, and
// deny-only groups are dropped when the identity is captured, so such a // deny-only groups are dropped when the identity is captured, so such a
// caller is correctly reported as unprivileged. Domain group memberships // caller is correctly reported as unprivileged. Domain group memberships
// (Domain Admins and friends) are deliberately not consulted: they say // (Domain Admins and friends) are deliberately not consulted: they say
@@ -83,8 +84,7 @@ func (i Identity) IsPrivileged() bool {
return true return true
} }
switch i.SID { if i.SID == sidLocalSystem {
case sidLocalSystem, sidLocalService, sidNetworkService:
return true return true
} }
@@ -64,3 +64,58 @@ func TestIdentitySameUser(t *testing.T) {
}) })
} }
} }
func TestIdentityIsPrivileged(t *testing.T) {
tests := []struct {
name string
id Identity
want bool
}{
{
name: "Root",
id: Identity{UID: 0, GID: 0},
want: true,
},
{
name: "Non-root",
id: Identity{UID: 1000, GID: 1000},
want: false,
},
{
name: "Local system windows",
id: Identity{SID: sidLocalSystem},
want: true,
},
{
name: "Windows elevated",
id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001", Elevated: true},
want: true,
},
{
name: "Admin group windows",
id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001", Groups: []string{sidAdministrators}},
want: true,
},
{
name: "Regular user windows",
id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001"},
want: false,
},
{
name: "Network service windows",
id: Identity{SID: sidNetworkService},
want: false,
},
{
name: "Local service windows",
id: Identity{SID: sidLocalService},
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
assert.Equal(t, tt.want, tt.id.IsPrivileged())
})
}
}
+9 -1
View File
@@ -45,7 +45,15 @@ func init() {
// matching there would let a non-elevated shell of an administrator account // matching there would let a non-elevated shell of an administrator account
// act as an administrator, which is the boundary the token check exists to // act as an administrator, which is the boundary the token check exists to
// keep. // keep.
selfMayDelegate = !id.IsPrivileged() selfMayDelegate = mayDelegate(id)
}
// mayDelegate reports whether a daemon running as id may extend its authority to
// callers sharing its identity. The shared service accounts are excluded: their
// SID is held by unrelated services, so matching on it would grant them the
// daemon's authority.
func mayDelegate(id Identity) bool {
return !id.IsPrivileged() && id.SID != sidLocalService && id.SID != sidNetworkService
} }
// IsDaemonSelf reports whether an identity is this very process. The JSON gateway // IsDaemonSelf reports whether an identity is this very process. The JSON gateway
+26 -1
View File
@@ -98,7 +98,7 @@ func TestIsPrivilegedCaller_SelfRule(t *testing.T) {
t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate }) t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate })
selfIdentity, selfKnown = tt.self, tt.selfKnown selfIdentity, selfKnown = tt.self, tt.selfKnown
selfMayDelegate = tt.selfKnown && !tt.self.IsPrivileged() selfMayDelegate = tt.selfKnown && mayDelegate(tt.self)
if got := IsPrivilegedCaller(tt.caller); got != tt.want { if got := IsPrivilegedCaller(tt.caller); got != tt.want {
t.Fatalf("IsPrivilegedCaller(%v) with daemon %v = %t, want %t", t.Fatalf("IsPrivilegedCaller(%v) with daemon %v = %t, want %t",
@@ -132,3 +132,28 @@ func TestIsPrivilegedCaller_ThisProcess(t *testing.T) {
t.Errorf("an unrelated identity %v was treated as privileged", other) t.Errorf("an unrelated identity %v was treated as privileged", other)
} }
} }
// The shared service accounts are held by unrelated services, so a daemon running
// as one of them must not extend its authority to every process with that SID.
func TestMayDelegate(t *testing.T) {
tests := []struct {
name string
self Identity
want bool
}{
{name: "unprivileged unix user", self: Identity{UID: 1000}, want: true},
{name: "root", self: Identity{UID: 0}, want: false},
{name: "unprivileged windows user", self: Identity{SID: "S-1-5-21-1-2-3-1001"}, want: true},
{name: "elevated windows user", self: Identity{SID: "S-1-5-21-1-2-3-1001", Elevated: true}, want: false},
{name: "local system", self: Identity{SID: sidLocalSystem}, want: false},
{name: "local service", self: Identity{SID: sidLocalService}, want: false},
{name: "network service", self: Identity{SID: sidNetworkService}, want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := mayDelegate(tt.self); got != tt.want {
t.Errorf("mayDelegate(%+v) = %v, want %v", tt.self, got, tt.want)
}
})
}
}
+41 -21
View File
@@ -135,9 +135,10 @@ type Conn struct {
// used to store the remote Rosenpass key for Relayed connection in case of connection update from ice // used to store the remote Rosenpass key for Relayed connection in case of connection update from ice
rosenpassRemoteKey []byte rosenpassRemoteKey []byte
wgProxyICE wgproxy.Proxy wgProxyICE wgproxy.Proxy
wgProxyRelay wgproxy.Proxy wgProxyRelay wgproxy.Proxy
handshaker *Handshaker relayedConnRef *relayClient.Conn
handshaker *Handshaker
guard *guard.Guard guard *guard.Guard
wg sync.WaitGroup wg sync.WaitGroup
@@ -560,7 +561,7 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
conn.mu.Lock() conn.mu.Lock()
defer conn.mu.Unlock() defer conn.mu.Unlock()
if conn.ctx.Err() != nil { if conn.ctx.Err() != nil || rci.relayedConn.Context().Err() != nil {
if err := rci.relayedConn.Close(); err != nil { if err := rci.relayedConn.Close(); err != nil {
conn.Log.Warnf("failed to close unnecessary relayed connection: %v", err) conn.Log.Warnf("failed to close unnecessary relayed connection: %v", err)
} }
@@ -575,7 +576,9 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
conn.Log.Errorf("failed to add relayed net.Conn to local proxy: %v", err) conn.Log.Errorf("failed to add relayed net.Conn to local proxy: %v", err)
return return
} }
wgProxy.SetDisconnectListener(conn.onRelayDisconnected) wgProxy.SetDisconnectListener(func() {
conn.onRelayDisconnected(rci.relayedConn)
})
conn.dumpState.NewLocalProxy() conn.dumpState.NewLocalProxy()
@@ -583,7 +586,7 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
if conn.isICEActive() { if conn.isICEActive() {
conn.Log.Debugf("do not switch to relay because current priority is: %s", conn.currentConnPriority.String()) conn.Log.Debugf("do not switch to relay because current priority is: %s", conn.currentConnPriority.String())
conn.setRelayedProxy(wgProxy) conn.setRelayedProxy(wgProxy, rci.relayedConn)
conn.statusRelay.SetConnected() conn.statusRelay.SetConnected()
conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, time.Now()) conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, time.Now())
return return
@@ -614,15 +617,26 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
conn.rosenpassRemoteKey = rci.rosenpassPubKey conn.rosenpassRemoteKey = rci.rosenpassPubKey
conn.currentConnPriority = conntype.Relay conn.currentConnPriority = conntype.Relay
conn.statusRelay.SetConnected() conn.statusRelay.SetConnected()
conn.setRelayedProxy(wgProxy) conn.setRelayedProxy(wgProxy, rci.relayedConn)
conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, updateTime) conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, updateTime)
conn.Log.Infof("start to communicate with peer via relay") conn.Log.Infof("start to communicate with peer via relay")
conn.doOnConnected(rci.rosenpassPubKey, rci.rosenpassAddr, updateTime) conn.doOnConnected(rci.rosenpassPubKey, rci.rosenpassAddr, updateTime)
} }
func (conn *Conn) onRelayDisconnected() { // onRelayDisconnected reports the teardown of a relayed connection. relayedConn
// names the connection the signal belongs to, so a signal that arrives after
// its connection was replaced is ignored instead of tearing down its successor.
// A nil relayedConn means the caller does not track generations and the current
// connection is always torn down.
func (conn *Conn) onRelayDisconnected(relayedConn *relayClient.Conn) {
conn.mu.Lock() conn.mu.Lock()
defer conn.mu.Unlock() defer conn.mu.Unlock()
if relayedConn != nil && conn.relayedConnRef != relayedConn {
conn.Log.Debugf("ignoring relay disconnect of a superseded connection")
return
}
conn.handleRelayDisconnectedLocked() conn.handleRelayDisconnectedLocked()
} }
@@ -646,6 +660,7 @@ func (conn *Conn) handleRelayDisconnectedLocked() {
_ = conn.wgProxyRelay.CloseConn() _ = conn.wgProxyRelay.CloseConn()
conn.wgProxyRelay = nil conn.wgProxyRelay = nil
} }
conn.relayedConnRef = nil
changed := conn.statusRelay.Get() != worker.StatusDisconnected changed := conn.statusRelay.Get() != worker.StatusDisconnected
if changed { if changed {
@@ -813,7 +828,8 @@ func (conn *Conn) evalStatus() ConnStatus {
// //
// The result is a tri-state: // The result is a tri-state:
// - ConnStatusConnected: all available transports are up // - ConnStatusConnected: all available transports are up
// - ConnStatusPartiallyConnected: relay is up but ICE is still pending/reconnecting // - ConnStatusPartiallyConnected: one transport carries the traffic and the other does
// not: relay up with ICE down, or ICE up with the shared relay transport down
// - ConnStatusDisconnected: no working transport // - ConnStatusDisconnected: no working transport
func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) { func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) {
defer func() { defer func() {
@@ -830,13 +846,14 @@ func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) {
} }
return evalConnStatus(connStatusInputs{ return evalConnStatus(connStatusInputs{
forceRelay: IsForceRelayed(), forceRelay: IsForceRelayed(),
peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(), peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(),
relayConnected: conn.statusRelay.Get() == worker.StatusConnected, relayConnected: conn.statusRelay.Get() == worker.StatusConnected,
remoteSupportsICE: conn.handshaker.RemoteICESupported(), relayTransportConnected: conn.workerRelay.IsTransportConnected(),
iceWorkerCreated: iceWorkerCreated, remoteSupportsICE: conn.handshaker.RemoteICESupported(),
iceStatusConnecting: conn.statusICE.Get() != worker.StatusDisconnected, iceWorkerCreated: iceWorkerCreated,
iceInProgress: iceInProgress, iceStatusConnected: conn.statusICE.Get() == worker.StatusConnected,
iceInProgress: iceInProgress,
}) })
} }
@@ -930,13 +947,14 @@ func (conn *Conn) logTraceConnState() {
} }
} }
func (conn *Conn) setRelayedProxy(proxy wgproxy.Proxy) { func (conn *Conn) setRelayedProxy(proxy wgproxy.Proxy, relayedConn *relayClient.Conn) {
if conn.wgProxyRelay != nil { if conn.wgProxyRelay != nil {
if err := conn.wgProxyRelay.CloseConn(); err != nil { if err := conn.wgProxyRelay.CloseConn(); err != nil {
conn.Log.Warnf("failed to close deprecated wg proxy conn: %v", err) conn.Log.Warnf("failed to close deprecated wg proxy conn: %v", err)
} }
} }
conn.wgProxyRelay = proxy conn.wgProxyRelay = proxy
conn.relayedConnRef = relayedConn
} }
// onWGHandshakeSuccess is called when the first WireGuard handshake is detected // onWGHandshakeSuccess is called when the first WireGuard handshake is detected
@@ -1044,19 +1062,21 @@ func evalConnStatus(in connStatusInputs) guard.ConnStatus {
return boolToConnStatus(relayUsedAndUp) return boolToConnStatus(relayUsedAndUp)
} }
// ICE counts as "up" when the status is anything other than Disconnected, OR // ICE counts as "running" when either connected or attempting to connect.
// when a negotiation is currently in progress (so we don't spam offers while one is in flight). iceRunning := in.iceStatusConnected || in.iceInProgress
iceUp := in.iceStatusConnecting || in.iceInProgress
// Relay side is acceptable if the peer doesn't rely on relay, or relay is connected. // Relay side is acceptable if the peer doesn't rely on relay, or relay is connected.
relayOK := !in.peerUsesRelay || in.relayConnected relayOK := !in.peerUsesRelay || in.relayConnected
switch { switch {
case iceUp && relayOK: case iceRunning && relayOK:
return guard.ConnStatusConnected return guard.ConnStatusConnected
case relayUsedAndUp: case relayUsedAndUp:
// Relay is up but ICE is down — partially connected. // Relay is up but ICE is down — partially connected.
return guard.ConnStatusPartiallyConnected return guard.ConnStatusPartiallyConnected
case in.iceStatusConnected && !in.relayTransportConnected:
// ICE is up and the shared relay transport is down — offers cannot restore it.
return guard.ConnStatusPartiallyConnected
default: default:
return guard.ConnStatusDisconnected return guard.ConnStatusDisconnected
} }
+8 -7
View File
@@ -17,13 +17,14 @@ const (
// tri-state connection classification. Extracted so the decision logic can be unit-tested // tri-state connection classification. Extracted so the decision logic can be unit-tested
// without constructing full Worker/Handshaker objects. // without constructing full Worker/Handshaker objects.
type connStatusInputs struct { type connStatusInputs struct {
forceRelay bool // NB_FORCE_RELAY or JS/WASM forceRelay bool // NB_FORCE_RELAY or JS/WASM
peerUsesRelay bool // remote peer advertises relay support AND local has relay peerUsesRelay bool // remote peer advertises relay support AND local has relay
relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay) relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay)
remoteSupportsICE bool // remote peer sent ICE credentials relayTransportConnected bool // the relay transport shared by all peers on that server is up
iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode) remoteSupportsICE bool // remote peer sent ICE credentials
iceStatusConnecting bool // statusICE is anything other than Disconnected iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode)
iceInProgress bool // a negotiation is currently in flight iceStatusConnected bool // statusICE reports Connected
iceInProgress bool // a negotiation is currently in flight
} }
// ConnStatus describe the status of a peer's connection // ConnStatus describe the status of a peer's connection
+71 -12
View File
@@ -30,6 +30,21 @@ func TestEvalConnStatus_ForceRelay(t *testing.T) {
}, },
want: guard.ConnStatusDisconnected, want: guard.ConnStatusDisconnected,
}, },
{
name: "force relay, relay up but the shared transport reports down",
in: connStatusInputs{
forceRelay: true,
peerUsesRelay: true,
relayConnected: true,
relayTransportConnected: false,
// The ICE inputs are set so that the force-relay return is the only branch
// that can produce Connected here: without it the peer would fall through to
// relayUsedAndUp and report PartiallyConnected.
remoteSupportsICE: true,
iceWorkerCreated: true,
},
want: guard.ConnStatusConnected,
},
{ {
name: "force relay, peer does NOT use relay - disconnected forever", name: "force relay, peer does NOT use relay - disconnected forever",
in: connStatusInputs{ in: connStatusInputs{
@@ -123,24 +138,28 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true in.peerUsesRelay = true
in.relayConnected = true in.relayConnected = true
in.iceStatusConnecting = true in.relayTransportConnected = true
in.iceStatusConnected = true
}, },
want: guard.ConnStatusConnected, want: guard.ConnStatusConnected,
}, },
{ {
name: "ICE connected, peer does NOT use relay", name: "ICE connected, peer does NOT use relay, shared transport down",
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false in.peerUsesRelay = false
in.relayConnected = false in.relayConnected = false
in.iceStatusConnecting = true in.relayTransportConnected = false
in.iceStatusConnected = true
}, },
// A peer that does not rely on relay is unaffected by the shared transport:
// relayOK is true, so the first arm matches before the transport is considered.
want: guard.ConnStatusConnected, want: guard.ConnStatusConnected,
}, },
{ {
name: "ICE InProgress only, peer does NOT use relay", name: "ICE InProgress only, peer does NOT use relay",
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false in.peerUsesRelay = false
in.iceStatusConnecting = false in.iceStatusConnected = false
in.iceInProgress = true in.iceInProgress = true
}, },
want: guard.ConnStatusConnected, want: guard.ConnStatusConnected,
@@ -150,7 +169,8 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true in.peerUsesRelay = true
in.relayConnected = true in.relayConnected = true
in.iceStatusConnecting = false in.relayTransportConnected = true
in.iceStatusConnected = false
in.iceInProgress = false in.iceInProgress = false
}, },
want: guard.ConnStatusPartiallyConnected, want: guard.ConnStatusPartiallyConnected,
@@ -160,21 +180,60 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false in.peerUsesRelay = false
in.relayConnected = false in.relayConnected = false
in.iceStatusConnecting = false in.iceStatusConnected = false
in.iceInProgress = false in.iceInProgress = false
}, },
want: guard.ConnStatusDisconnected, want: guard.ConnStatusDisconnected,
}, },
{ {
name: "ICE up, peer uses relay but relay down -> partial (relay required, ICE ignored)", name: "ICE connected, relay down for this peer but the shared transport is up -> disconnected",
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true in.peerUsesRelay = true
in.relayConnected = false in.relayConnected = false
in.iceStatusConnecting = true in.relayTransportConnected = true
in.iceStatusConnected = true
},
// The transport is fine, so the peer itself is unreachable over relay: it may have
// moved to another server, and only an offer carries its new relay address.
want: guard.ConnStatusDisconnected,
},
{
name: "ICE connected, the shared relay transport is down -> partial",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = false
in.relayTransportConnected = false
in.iceStatusConnected = true
},
// ICE carries the traffic and the relay transport is restored by the relay client's
// own guard, not by offers, so this must not trigger the aggressive retry.
want: guard.ConnStatusPartiallyConnected,
},
{
name: "ICE only negotiating while the shared relay transport is down -> disconnected",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = false
in.relayTransportConnected = false
in.iceStatusConnected = false
in.iceInProgress = true
},
// A negotiation in flight is not a working transport, so this peer has no path at
// all and must keep the aggressive retry. Calling it partially connected spends the
// ICE retry budget and parks the guard on the hourly ticker, and nothing wakes it
// when the negotiation then fails: onICEStateDisconnected is only reached once ICE
// has reached Connected (worker_ice.go onConnectionStateChange).
want: guard.ConnStatusDisconnected,
},
{
name: "ICE down and the shared relay transport is down -> disconnected",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = false
in.relayTransportConnected = false
in.iceStatusConnected = false
in.iceInProgress = false
}, },
// relayOK = false (peer uses relay but it's down), iceUp = true
// first switch arm fails (relayOK false), relayUsedAndUp = false (relay down),
// falls into default: Disconnected.
want: guard.ConnStatusDisconnected, want: guard.ConnStatusDisconnected,
}, },
{ {
@@ -182,7 +241,7 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false in.peerUsesRelay = false
in.relayConnected = true // not actually used since peer doesn't rely on it in.relayConnected = true // not actually used since peer doesn't rely on it
in.iceStatusConnecting = false in.iceStatusConnected = false
in.iceInProgress = false in.iceInProgress = false
}, },
want: guard.ConnStatusDisconnected, want: guard.ConnStatusDisconnected,

Some files were not shown because too many files have changed in this diff Show More