Merge remote-tracking branch 'origin/main' into jnfrati/ubi-signal

This commit is contained in:
jnfrati
2026-10-09 11:57:41 +02:00
751 changed files with 58846 additions and 25144 deletions
+10
View File
@@ -14,5 +14,15 @@ reviews:
- "!**/*.ts"
- "!**/*.js"
- "!**/*.svg"
pre_merge_checks:
custom_checks:
- name: "No attribution trailers"
mode: error
instructions: >-
Fail when the PR description or any commit message carries an
attribution trailer or footer: Co-Authored-By, Claude-Session,
Generated-By, or a "Generated with"/"Generated by" tool line.
Contributors own their contributions (AGENTS.md); ask for the
lines to be removed.
chat:
auto_reply: true
+26
View File
@@ -0,0 +1,26 @@
#!/bin/bash
# Refuses commit messages that carry attribution trailers. Contributors own
# their contributions (AGENTS.md, "No Co-Authored-By or tool-attribution
# trailers"); a trailer spreads that ownership onto a tool or a bystander.
msg_file="$1"
# Trailer keys in any casing, with any bullet or emoji in front.
trailers='^[^[:alnum:]]*(co-authored-by|claude-session|generated-by):'
# "Generated with/by" footers, including "Generated with <emoji> by".
footer='^[^[:alnum:]]*generated (with|by)( [^[:alnum:]]*by)? '
# A footer names a product, so a capitalized word must follow the phrase
# itself. Prose such as "generated by the protobuf compiler" stays legal.
tool='[Gg][Ee][Nn][Ee][Rr][Aa][Tt][Ee][Dd] ([Ww][Ii][Tt][Hh]|[Bb][Yy])( [^[:alnum:]]*[Bb][Yy])? [^[:alnum:]]*[A-Z]'
offending=$( {
grep -Ein "$trailers" "$msg_file"
grep -Ein "$footer" "$msg_file" | grep -E "$tool"
} | sort -un )
if [ -n "$offending" ]; then
echo "commit-msg: attribution trailers are not accepted in this repository:" >&2
printf '%s\n' "$offending" | sed 's/^/ /' >&2
echo "Remove them and commit again (see AGENTS.md)." >&2
exit 1
fi
+22
View File
@@ -46,3 +46,25 @@ updates:
wireguard:
patterns:
- "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::"
@@ -34,7 +34,7 @@ jobs:
while IFS= read -r dir; do
echo "=== Checking $dir ==="
# Search for problematic imports, excluding test files
RESULTS=$(grep -r "github.com/netbirdio/netbird/\(management\|signal\|relay\|proxy\)" "$dir" --include="*.go" 2>/dev/null | grep -v "_test.go" | grep -v "test_" | grep -v "/test/" | grep -v "tools/idp-migrate/" || true)
RESULTS=$(grep -r "github.com/netbirdio/netbird/\(management\|signal\|relay\|proxy\)" "$dir" --include="*.go" 2>/dev/null | grep -v "_test.go" | grep -v "test_" | grep -v "/test/" | grep -v "tools/idp-migrate/" | grep -v "tools/mysql-migrate/" || true)
if [ -n "$RESULTS" ]; then
echo "❌ Found problematic dependencies:"
echo "$RESULTS"
@@ -93,7 +93,7 @@ jobs:
IMPORTERS=$(go list -json -deps ./... 2>/dev/null | jq -r "select(.Imports[]? == \"$package\") | .ImportPath")
# Check if any importer is NOT in management/signal/relay
BSD_IMPORTER=$(echo "$IMPORTERS" | grep -v "github.com/netbirdio/netbird/\(management\|signal\|relay\|proxy\|combined\|tools/idp-migrate\)" | head -1)
BSD_IMPORTER=$(echo "$IMPORTERS" | grep -v "github.com/netbirdio/netbird/\(management\|signal\|relay\|proxy\|combined\|tools/idp-migrate\|tools/mysql-migrate\)" | head -1)
if [ -n "$BSD_IMPORTER" ]; then
echo "❌ $package ($license) is imported by BSD-licensed code: $BSD_IMPORTER"
+2
View File
@@ -12,6 +12,8 @@ jobs:
docs-ack:
name: Require docs PR URL or explicit "not needed"
runs-on: ubuntu-latest
# Crowdin's translation-sync service PRs are auto-generated without the PR template.
if: github.event.pull_request.user.login != 'netbirddev'
steps:
- name: Read PR body
+3 -3
View File
@@ -38,12 +38,12 @@ jobs:
persist-credentials: false
- name: Set up Node.js
uses: actions/setup-node@v4
uses: actions/setup-node@v7
with:
node-version: "22"
- name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with:
version: 11
@@ -79,7 +79,7 @@ jobs:
run: echo "path=$(pnpm store path --silent)" >> "$GITHUB_OUTPUT"
- name: Cache pnpm store
uses: actions/cache@v4
uses: actions/cache@v6
with:
path: ${{ steps.pnpm-store.outputs.path }}
key: ${{ runner.os }}-pnpm-${{ hashFiles('client/ui/frontend/pnpm-lock.yaml') }}
+7 -5
View File
@@ -46,15 +46,17 @@ jobs:
run: git --no-pager diff --exit-code
- name: Test
# Exclude client/ui: its main.go uses //go:embed all:frontend/dist,
# which fails to compile until the frontend has been built. The Wails UI
# has no Go-side unit tests, and its release pipeline runs `pnpm build`
# before goreleaser.
# Exclude the client/ui package itself: its main.go uses //go:embed
# all:frontend/dist, which fails to compile until the frontend has been
# built, and its release pipeline runs `pnpm build` before goreleaser.
# The pattern is anchored so the subpackages (services, preferences,
# i18n, authsession) still run: they hold Go-side unit tests and need no
# frontend bundle.
# `go list -e` lets the listing succeed even though the embed fails to
# resolve; the grep then drops the broken package by path. Without -e,
# go list aborts with empty stdout and `go test` falls back to the repo
# root, which has no Go files.
run: NETBIRD_STORE_ENGINE=${{ matrix.store }} CI=true go test -coverprofile=coverage.txt -tags 'devcert privileged' -exec 'sudo --preserve-env=CI,NETBIRD_STORE_ENGINE' -timeout 5m -p 1 $(go list -e ./... | grep -v -e /management -e /signal -e /relay -e /proxy -e /combined -e /client/ui -e /client/testutil/privileged)
run: NETBIRD_STORE_ENGINE=${{ matrix.store }} CI=true go test -coverprofile=coverage.txt -tags 'devcert privileged' -exec 'sudo --preserve-env=CI,NETBIRD_STORE_ENGINE' -timeout 5m -p 1 $(go list -e ./... | grep -v -e /management -e /signal -e /relay -e /proxy -e /combined -e '/client/ui$' -e /client/testutil/privileged)
- name: Upload coverage reports to Codecov
uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f #v7.0.0
+1 -1
View File
@@ -33,7 +33,7 @@ jobs:
with:
usesh: true
copyback: false
release: "15.0"
release: "15.1"
envs: "GO_VERSION"
prepare: |
pkg install -y curl pkgconf xorg
+73 -7
View File
@@ -160,9 +160,10 @@ jobs:
- name: Test
# Exclude client/ui: its main.go uses //go:embed all:frontend/dist,
# which fails to compile until the frontend has been built. The Wails UI
# has no Go-side unit tests, and its release pipeline runs `pnpm build`
# before goreleaser.
# which fails to compile until the frontend has been built, and its
# release pipeline runs `pnpm build` before goreleaser. The subpackages
# go with it because this runner's gtk4 is older than the wails runtime
# needs; the Client UI / Unit job below covers them instead.
# `go list -e` lets the listing succeed even though the embed fails to
# resolve; the grep then drops the broken package by path. Without -e,
# go list aborts with empty stdout and `go test` falls back to the repo
@@ -177,6 +178,35 @@ jobs:
slug: netbirdio/netbird
flags: unit,client
test_client_ui:
name: "Client UI / Unit"
# Pinned to 24.04 rather than the 22.04 the other client jobs use: the wails
# runtime's linux cgo layer needs GtkFileDialog, which arrived in gtk4 4.10,
# and jammy ships 4.6. Not ubuntu-latest, so a runner image rollover cannot
# move this out from under us.
runs-on: ubuntu-24.04
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: Install dependencies
run: sudo apt update && sudo apt install -y -q libgtk-4-dev libwebkitgtk-6.0-dev libsoup-3.0-dev
- name: Test
# client/ui itself stays out: its main.go embeds all:frontend/dist,
# which only exists after `pnpm build`. The subpackages carry the
# Go-side unit tests, including the window manager re-entrancy
# regression test, and need no frontend bundle.
run: CGO_ENABLED=1 go test -timeout 5m ./client/ui/authsession/... ./client/ui/i18n/... ./client/ui/preferences/... ./client/ui/services/...
test_client_on_docker:
name: "Client (Docker) / Unit"
needs: [build-cache]
@@ -211,6 +241,9 @@ jobs:
${{ runner.os }}-gotest-cache-
- name: Run tests in container
# Unlike the native job above, this one drops all of client/ui including
# the subpackages: the alpine container has no gtk4/webkitgtk, so the
# Wails application package they import would fail to link.
env:
HOST_GOCACHE: ${{ steps.go-env.outputs.cache_dir }}
HOST_GOMODCACHE: ${{ steps.go-env.outputs.modcache_dir }}
@@ -237,7 +270,7 @@ jobs:
sh -c ' \
apk update; apk add --no-cache \
ca-certificates iptables ip6tables dbus dbus-dev libpcap-dev build-base; \
go test -buildvcs=false -tags "devcert privileged" -v -timeout 10m -p 1 $(go list -e -buildvcs=false ./... | grep -v -e /management -e /signal -e /relay -e /proxy -e /combined -e /client/ui -e /upload-server -e /client/testutil/privileged)
go test -buildvcs=false -tags "devcert privileged" -v -timeout 10m -p 1 $(go list -e -buildvcs=false ./... | grep -v -e /management -e /signal -e /relay -e /proxy -e /combined -e /client/ui -e /upload-server -e /client/testutil/privileged -e tools/mysql-migrate)
'
test_relay:
@@ -481,14 +514,32 @@ jobs:
if: matrix.store == 'mysql'
run: docker pull mlsmaycon/warmed-mysql:8
# The -json stream goes through tools/gotestsummary so the log shows one
# line per test, the output of failed tests, the head of a timeout panic
# with the still-running tests, and the slowest tests per package.
- name: Test
shell: bash
run: |
set -o pipefail
CGO_ENABLED=1 GOARCH=${{ matrix.arch }} \
NETBIRD_STORE_ENGINE=${{ matrix.store }} \
CI=true \
go test -tags=devcert -coverprofile=coverage.txt \
go test -json -tags=devcert -coverprofile=coverage.txt \
-exec "sudo --preserve-env=CI,NETBIRD_STORE_ENGINE" \
-timeout 20m ./management/... ./shared/management/...
-timeout 20m ./management/... ./shared/management/... \
| tee management-test-events.jsonl \
| go run ./tools/gotestsummary
# The summary trims long outputs; the raw stream keeps every line for
# the failures that need it. A green run has no use for it.
- name: Upload raw test events
if: failure()
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a #v7.0.1
with:
name: management-unit-test-events-${{ matrix.store }}
path: management-test-events.jsonl
if-no-files-found: ignore
retention-days: 14
- name: Upload coverage reports to Codecov
if: matrix.arch == 'amd64'
@@ -738,12 +789,27 @@ jobs:
- name: check git status
run: git --no-pager diff --exit-code
# Same summary as the unit job: a timeout here names the tests still
# running instead of ending in a goroutine dump.
- name: Test
shell: bash
run: |
set -o pipefail
CGO_ENABLED=1 GOARCH=${{ matrix.arch }} \
NETBIRD_STORE_ENGINE=${{ matrix.store }} \
CI=true \
mage integrationtest:all -gotestflags="-coverprofile=coverage.txt"
mage integrationtest:all -gotestflags="-json -coverprofile=coverage.txt" \
| tee management-integration-test-events.jsonl \
| go run ./tools/gotestsummary
- name: Upload raw test events
if: failure()
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a #v7.0.1
with:
name: management-integration-test-events-${{ matrix.store }}
path: management-integration-test-events.jsonl
if-no-files-found: ignore
retention-days: 14
- name: Upload coverage reports to Codecov
if: matrix.arch == 'amd64'
+7 -5
View File
@@ -66,15 +66,17 @@ jobs:
- run: PsExec64 -s -w ${{ github.workspace }} C:\hostedtoolcache\windows\go\${{ steps.go.outputs.go-version }}\x64\bin\go.exe env -w GOCACHE=${{ env.modcache }}
- run: PsExec64 -s -w ${{ github.workspace }} C:\hostedtoolcache\windows\go\${{ steps.go.outputs.go-version }}\x64\bin\go.exe mod tidy
- name: Generate test script
# Exclude client/ui: its main.go uses //go:embed all:frontend/dist,
# which fails to compile until the frontend has been built. The Wails UI
# has no Go-side unit tests, and its release pipeline runs `pnpm build`
# before goreleaser.
# Exclude the client/ui package itself: its main.go uses //go:embed
# all:frontend/dist, which fails to compile until the frontend has been
# built, and its release pipeline runs `pnpm build` before goreleaser.
# The pattern is anchored so the subpackages (services, preferences,
# i18n, authsession) still run: they hold Go-side unit tests and need no
# frontend bundle.
# `go list -e` lets the listing succeed even though the embed fails to
# resolve; the Where-Object pipeline then drops the broken package by
# path. Without -e, go list aborts with empty stdout.
run: |
$packages = go list -e ./... | Where-Object { $_ -notmatch '/management' } | Where-Object { $_ -notmatch '/relay' } | Where-Object { $_ -notmatch '/signal' } | Where-Object { $_ -notmatch '/proxy' } | Where-Object { $_ -notmatch '/combined' } | Where-Object { $_ -notmatch '/client/ui' }
$packages = go list -e ./... | Where-Object { $_ -notmatch '/management' } | Where-Object { $_ -notmatch '/relay' } | Where-Object { $_ -notmatch '/signal' } | Where-Object { $_ -notmatch '/proxy' } | Where-Object { $_ -notmatch '/combined' } | Where-Object { $_ -notmatch '/client/ui$' } | Where-Object { $_ -notmatch '/tools/mysql-migrate' }
$goExe = "C:\hostedtoolcache\windows\go\${{ steps.go.outputs.go-version }}\x64\bin\go.exe"
$cmd = "$goExe test -tags `"devcert privileged`" -timeout 10m -p 1 $($packages -join ' ') > test-out.txt 2>&1"
Set-Content -Path "${{ github.workspace }}\run-tests.cmd" -Value $cmd
+47 -1
View File
@@ -30,7 +30,7 @@ jobs:
# segment by codespell and behave the same across versions; the
# recursive "**" form did not take effect with the codespell shipped
# by this action.
skip: go.mod,go.sum,*/proxy/web/*,*pnpm-lock.yaml,*package-lock.json,*/locales/de/*,*/locales/es/*,*/locales/fr/*,*/locales/hu/*,*/locales/it/*,*/locales/pt/*,*/locales/ru/*,*/locales/zh-CN/*,*/i18n/TRANSLATING.md
skip: go.mod,go.sum,*/proxy/web/*,*pnpm-lock.yaml,*package-lock.json,*/locales/de/*,*/locales/es/*,*/locales/fr/*,*/locales/gl/*,*/locales/hu/*,*/locales/it/*,*/locales/pt/*,*/locales/ru/*,*/locales/zh-CN/*,*/i18n/TRANSLATING.md
golangci:
strategy:
fail-fast: false
@@ -80,3 +80,49 @@ jobs:
skip-save-cache: true
cache-invalidation-interval: 0
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/...
+2
View File
@@ -7,6 +7,8 @@ on:
jobs:
check-title:
runs-on: ubuntu-latest
# Crowdin's translation-sync service PRs are auto-generated with a fixed title.
if: github.event.pull_request.user.login != 'netbirddev'
steps:
- name: Validate PR title prefix
uses: actions/github-script@3a2844b7e9c422d3c10d287c895573f7108da1b3 # v9.0.0
+201
View File
@@ -0,0 +1,201 @@
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
- netbird-server
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"
"netbird-server ghcr.io/netbirdio/netbird-server -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 globstar
results=(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
+57 -12
View File
@@ -69,7 +69,7 @@ jobs:
with:
usesh: true
copyback: false
release: "15.0"
release: "15.1"
envs: "GO_VERSION"
prepare: |
# Install required packages
@@ -186,6 +186,22 @@ jobs:
run: bash shared/management/http/api/generate.sh
- name: check git status
run: git --no-pager diff --exit-code
- name: Generate RPM changelog from git tags
# nfpm embeds changelog.yml into the RPM; Red Hat software certification
# requires a changelog. Generated, not committed (see .gitignore).
# chglog is a go.mod tool directive, so go.sum pins it and its deps.
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
# release_files/collect-licenses.sh -w reads the proxy UI's license terms from node_modules.
working-directory: proxy/web
run: npm ci --ignore-scripts
- name: Set up QEMU
uses: docker/setup-qemu-action@06116385d9baf250c9f4dcb4858b16962ea869c3 #v4.1.0
- name: Set up Docker Buildx
@@ -225,14 +241,18 @@ jobs:
uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2
with:
version: ${{ env.GORELEASER_VER }}
args: release --clean ${{ env.flags }}
args: release --config .goreleaser.generated.yaml --clean ${{ env.flags }}
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
HOMEBREW_TAP_GITHUB_TOKEN: ${{ secrets.HOMEBREW_TAP_GITHUB_TOKEN }}
UPLOAD_DEBIAN_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }}
UPLOAD_YUM_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }}
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_DOCKER_PUSH: ${{ env.SKIP_DOCKER_PUSH }}
- name: Verify RPM signatures
@@ -287,10 +307,17 @@ jobs:
image_refs=()
tag_and_push() {
local src="$1" img_name tag dst
local src="$1" img_name tag dst variant=""
img_name="${src%%:*}"
# Variants share a repository with their default image, so keep
# their tag suffixes. Order matters: the first matching pattern wins.
case "$src" in
*-rootless-ubi-amd64) variant="-rootless-ubi" ;;
*-rootless-amd64) variant="-rootless" ;;
*-ubi-amd64) variant="-ubi" ;;
esac
for tag in $(resolve_tags); do
dst="${img_name}:${tag}"
dst="${img_name}:${tag}${variant}"
echo "Tagging ${src} -> ${dst}"
docker tag "$src" "$dst"
docker push "$dst"
@@ -353,6 +380,24 @@ jobs:
path: dist/netbird_darwin**
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:
runs-on: ubuntu-latest
outputs:
@@ -407,12 +452,12 @@ jobs:
run: git --no-pager diff --exit-code
- name: Set up Node.js
uses: actions/setup-node@v4
uses: actions/setup-node@v7
with:
node-version: '22'
- name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with:
version: 11
@@ -544,12 +589,12 @@ jobs:
run: git --no-pager diff --exit-code
- name: Set up Node.js
uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4
uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7.0.0
with:
node-version: '22'
- name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with:
version: 11
@@ -641,11 +686,11 @@ jobs:
- name: check git status
run: git --no-pager diff --exit-code
- name: Set up Node.js
uses: actions/setup-node@v4
uses: actions/setup-node@v7
with:
node-version: '22'
- name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with:
version: 11
- name: Install wails3 CLI
@@ -764,7 +809,7 @@ jobs:
run: 7z x -o"${{ github.workspace }}/NSIS_Plugins" "${{ github.workspace }}/ShellExecAsUser_amd64-Unicode.7z"
- name: Set up Go for wails3 CLI
uses: actions/setup-go@v5
uses: actions/setup-go@v6
with:
go-version-file: "go.mod"
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
+4 -3
View File
@@ -32,11 +32,12 @@ jobs:
persist-credentials: false
- name: Set up Node.js
uses: actions/setup-node@v4
uses: actions/setup-node@v7
with:
node-version: "22"
# English (en) is the source of truth for translation keys; every other
# locale declared in _index.json must carry the exact same key set.
# English (en) is the source of truth for translation keys. Locales declared
# in _index.json fail on orphaned keys or placeholder mismatches; missing
# keys only warn, since they fall back to English at runtime.
- name: Check translation key parity
run: node client/ui/i18n/check-translations.mjs
+7
View File
@@ -35,3 +35,10 @@ vendor/
/netbird
client/netbird-electron/
management/server/types/testdata/
# generated by chglog in the release workflow, embedded into the RPM
changelog.yml
# generated by rpm-provides.sh, the config GoReleaser actually runs
.goreleaser.generated.yaml
.chglog.yml
+217 -6
View File
@@ -40,6 +40,32 @@ builds:
tags:
- load_wgnt_from_rsrc
# Single-arch builds: nfpm provides is not templated, so the RPM splits per arch.
- &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
- <<: *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
dir: client
binary: netbird
@@ -190,6 +216,28 @@ builds:
- -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 }}"
- id: netbird-mysql-migrate
dir: tools/mysql-migrate
env:
- CGO_ENABLED=1
- >-
{{- if eq .Runtime.Goos "linux" }}
{{- if eq .Arch "arm64"}}CC=aarch64-linux-gnu-gcc{{- end }}
{{- if eq .Arch "arm"}}CC=arm-linux-gnueabihf-gcc{{- end }}
{{- end }}
binary: netbird-mysql-migrate
goos:
- linux
goarch:
- amd64
- arm64
- arm
goarm:
- 7
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 }}"
universal_binaries:
- id: netbird
@@ -206,6 +254,10 @@ archives:
builds:
- netbird-idp-migrate
name_template: "netbird-idp-migrate_{{ .Version }}_{{ .Os }}_{{ .Arch }}"
- id: netbird-mysql-migrate
builds:
- netbird-mysql-migrate
name_template: "netbird-mysql-migrate_{{ .Version }}_{{ .Os }}_{{ .Arch }}"
nfpms:
- maintainer: Netbird <dev@netbird.io>
@@ -223,23 +275,72 @@ nfpms:
postinstall: "release_files/post_install.sh"
preremove: "release_files/pre_remove.sh"
- maintainer: Netbird <dev@netbird.io>
- &netbird_rpm
maintainer: Netbird <dev@netbird.io>
description: Netbird client.
homepage: https://netbird.io/
license: BSD-3-Clause
vendor: NetBird
id: netbird_rpm
id: netbird_rpm_amd64
bindir: /usr/bin
builds:
- netbird
ids:
- netbird-rpm-amd64
formats:
- 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
# store. Red Hat software certification (RPM Dependency Tracking) also
# rejects packages that declare no dependencies at all.
dependencies:
- ca-certificates
# Generated in CI by chglog from git tags; Red Hat certification requires an
# RPM changelog (RPM Version Handling subtest).
changelog: changelog.yml
# License, documentation and a config file so the RPM Provenance subtest sees
# %license, %doc and %config entries instead of a bare binary.
contents:
- src: LICENSE
dst: /usr/share/licenses/netbird/LICENSE
type: license
- src: README.md
dst: /usr/share/doc/netbird/README.md
type: doc
- src: release_files/netbird.sysconfig
dst: /etc/sysconfig/netbird
type: config|noreplace
scripts:
postinstall: "release_files/post_install.sh"
preremove: "release_files/pre_remove.sh"
rpm:
summary: NetBird client
group: Applications/Internet
packager: NetBird <dev@netbird.io>
signature:
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:
- id: netbird
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
@@ -289,6 +390,43 @@ dockers_v2:
"org.opencontainers.image.revision": "{{.FullCommit}}"
"org.opencontainers.image.source": "{{.GitURL}}"
"maintainer": "dev@netbird.io"
- id: netbird-rootless-ubi
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
ids:
- netbird
images:
- netbirdio/netbird
- ghcr.io/netbirdio/netbird
tags:
- "{{ .Version }}-rootless-ubi"
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}rootless-ubi-latest{{ end }}"
dockerfile: client/Dockerfile-rootless.ubi
extra_files:
- client/netbird-entrypoint.sh
platforms:
- linux/amd64
- linux/arm64
build_args:
VERSION: "{{ .Version }}"
RELEASE: "{{ .Timestamp }}"
hooks:
pre:
- cmd: 'sh release_files/collect-licenses.sh -t load_wgnt_from_rsrc "{{ .ContextDir }}/licenses" ./client 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"
- id: relay
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
ids:
@@ -400,7 +538,7 @@ dockers_v2:
tags:
- "{{ .Version }}"
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}latest{{ end }}"
dockerfile: upload-server/Dockerfile
dockerfile: upload-server/Dockerfile.release
platforms:
- linux/amd64
- linux/arm64
@@ -434,6 +572,41 @@ dockers_v2:
"org.opencontainers.image.revision": "{{.FullCommit}}"
"org.opencontainers.image.source": "{{.GitURL}}"
"maintainer": "dev@netbird.io"
- id: netbird-server-ubi
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
ids:
- netbird-server
images:
- netbirdio/netbird-server
- ghcr.io/netbirdio/netbird-server
tags:
- "{{ .Version }}-ubi"
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}ubi-latest{{ end }}"
dockerfile: combined/Dockerfile.ubi
platforms:
- linux/amd64
- linux/arm64
build_args:
VERSION: "{{ .Version }}"
RELEASE: "{{ .Timestamp }}"
hooks:
pre:
- cmd: 'sh release_files/collect-licenses.sh -l combined/LICENSE "{{ .ContextDir }}/licenses" ./combined amd64 arm64'
env:
- GOOS=linux
- CGO_ENABLED=1
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"
- id: netbird-proxy
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
ids:
@@ -456,6 +629,41 @@ dockers_v2:
"org.opencontainers.image.revision": "{{.FullCommit}}"
"org.opencontainers.image.source": "{{.GitURL}}"
"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 release_files/collect-licenses.sh -l proxy/LICENSE -w "{{ .ContextDir }}/licenses" ./proxy/cmd/proxy 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:
- ids:
@@ -488,7 +696,10 @@ uploads:
- name: yum
skip: "{{ .Env.SKIP_PUBLISH }}"
ids:
- netbird_rpm
- netbird_rpm_amd64
- netbird_rpm_arm64
- netbird_rpm_arm
- netbird_rpm_386
mode: archive
target: https://pkgs.wiretrustee.com/yum/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }}
username: dev@wiretrustee.com
+1 -1
View File
@@ -77,7 +77,7 @@ make lint # golangci-lint on files changed vs origin/main (also the p
make lint-all # full-repository lint, matches CI
make test-unit # host-safe unit tests, -tags devcert, no sudo
make test-privileged # privileged-tagged suite in a Docker container with NET_ADMIN
make setup-hooks # wire make lint into .githooks/pre-push
make setup-hooks # wire .githooks: pre-push runs make lint, commit-msg refuses attribution trailers
# Narrow runs
go test ./client/internal/dns/...
+4 -1
View File
@@ -1 +1,4 @@
See [AGENTS.md](AGENTS.md) for the agent guidelines in this repository.
The agent guidelines live in [AGENTS.md](AGENTS.md). It is imported here so
every session loads it in full rather than following a pointer.
@AGENTS.md
+1 -1
View File
@@ -1,4 +1,4 @@
This BSD‑3‑Clause license applies to all parts of the repository except for the directories management/, signal/, relay/ and combined/.
This BSD-3-Clause license applies to all parts of the repository except for the directories management/, signal/, relay/ and combined/.
Those directories are licensed under the GNU Affero General Public License version 3.0 (AGPLv3). See the respective LICENSE files inside each directory.
BSD 3-Clause License
+2 -2
View File
@@ -23,8 +23,8 @@ lint-install: $(GOLANGCI_LINT)
# Setup git hooks for all developers
setup-hooks:
@git config core.hooksPath .githooks
@chmod +x .githooks/pre-push
@echo "✅ Git hooks configured! Pre-push will now run 'make lint'"
@chmod +x .githooks/pre-push .githooks/commit-msg
@echo "✅ Git hooks configured! Pre-push runs 'make lint'; commit-msg refuses attribution trailers"
# Host-safe unit tests: excludes the privileged-tagged tests (root / system-mutating).
# Runs as a normal user with no sudo and leaves host networking untouched.
+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.
### 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
- [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
+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).
Nothing else in the account.
- **`usage_viewer`** — the regular User baseline plus read on
`agent_network.usage` (the aggregated usage and cost overview) and read-only
access to the resources the usage filters resolve against: users, groups,
peers, and the provider list (connection config redacted — no upstream URLs
or operator-supplied header values). No policies, and no account-wide
request-level access logs; like any caller, it still reads its own requests
through the self-scoped endpoints below.
`agent_network.usage` (the aggregated usage and cost overview) and
`agent_network.logs` (the account-wide request-level access logs, which can
contain captured prompts), and read-only access to the resources those
filters resolve against: users, groups, peers, and the provider list
(connection config redacted — no upstream URLs or operator-supplied header
values). No policies, guardrails, budgets, or settings.
Every authenticated user, regardless of role, can read the caller-scoped
self-service endpoint `GET /api/agent-network/agent-config` (the endpoint, providers,
+52 -33
View File
@@ -3,56 +3,75 @@ package base62
import (
"fmt"
"math"
"strings"
)
const (
alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
base = uint32(len(alphabet))
alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
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.
func Encode(num uint32) string {
if num == 0 {
return string(alphabet[0])
// The returned string will be between 1-6 characters long.
func Encode(n uint32) string {
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
for num > 0 {
remainder := num % base
encoded.WriteByte(alphabet[remainder])
num /= base
}
// Reverse the encoded string
encodedString := encoded.String()
reversed := reverse(encodedString)
return reversed
return string(buf[idx:])
}
// 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) {
if len(encoded) == 0 {
return 0, ErrEmptyString
}
var decoded uint32
strLen := len(encoded)
for i, char := range encoded {
index := strings.IndexRune(alphabet, char)
for _, char := range encoded {
index := int8(-1)
if int(char) < len(charToIndex) {
index = charToIndex[char]
}
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
}
// 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
import (
"errors"
"math"
"testing"
)
func TestEncodeDecode(t *testing.T) {
tests := []struct {
num uint32
testCases := []struct {
input uint32
expected string
}{
{0},
{1},
{42},
{12345},
{99999},
{123456789},
{0, "0"},
{1, "1"},
{5, "5"},
{9, "9"},
{10, "A"},
{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 {
encoded := Encode(tt.num)
for _, tc := range testCases {
encoded := Encode(tc.input)
if encoded != tc.expected {
t.Errorf("Encode(%d) = %s; want %s", tc.input, encoded, tc.expected)
}
decoded, err := Decode(encoded)
if err != nil {
t.Errorf("Decode error: %v", err)
t.Errorf("Expected error nil, got %v", err)
}
if decoded != tt.num {
t.Errorf("Decode(%v) = %v, want %v", encoded, decoded, tt.num)
if decoded != tc.input {
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)
}
}
+45
View File
@@ -0,0 +1,45 @@
FROM registry.access.redhat.com/ubi9/ubi-minimal@sha256:7fbeae18dc9476399f565e68255f602a3374ea8614ba3d14843565131a13ff93
ARG TARGETPLATFORM
ARG NETBIRD_BINARY=$TARGETPLATFORM/netbird
ARG VERSION=dev
ARG RELEASE=1
LABEL name="netbird-rootless" \
maintainer="NetBird <dev@netbird.io>" \
vendor="NetBird GmbH" \
version="${VERSION}" \
release="${RELEASE}" \
summary="NetBird Rootless Client" \
description="NetBird connects devices through an encrypted overlay using userspace networking without a TUN device or network administration capabilities."
RUN microdnf install -y bash ca-certificates && microdnf clean all
COPY --chmod=0555 client/netbird-entrypoint.sh /usr/local/bin/netbird-entrypoint.sh
COPY --chmod=0555 ${NETBIRD_BINARY} /usr/local/bin/netbird
COPY licenses/ /licenses/
# Only application storage is group-writable for arbitrary non-root UIDs.
# Runtime-created credentials keep the client's restrictive file modes.
RUN mkdir -p /var/lib/netbird && \
chown 1000:0 /var/lib/netbird && \
chmod 0770 /var/lib/netbird && \
chmod -R a+rX /licenses
WORKDIR /var/lib/netbird
USER 1000:0
ENV \
HOME="/var/lib/netbird" \
NETBIRD_BIN="/usr/local/bin/netbird" \
NB_USE_NETSTACK_MODE="true" \
NB_ENABLE_NETSTACK_LOCAL_FORWARDING="true" \
NB_CONFIG="/var/lib/netbird/config.json" \
NB_STATE_DIR="/var/lib/netbird" \
NB_DAEMON_ADDR="unix:///var/lib/netbird/netbird.sock" \
NB_LOG_FILE="console,/var/lib/netbird/client.log" \
NB_DISABLE_DNS="true" \
NB_ENABLE_CAPTURE="false" \
NB_ENTRYPOINT_SERVICE_TIMEOUT="30"
STOPSIGNAL SIGTERM
ENTRYPOINT ["/usr/local/bin/netbird-entrypoint.sh"]
+43 -11
View File
@@ -9,6 +9,7 @@ import (
"slices"
"strings"
"sync"
"sync/atomic"
"time"
"golang.org/x/exp/maps"
@@ -90,13 +91,20 @@ type Client struct {
connectClient *internal.ConnectClient
config *profilemanager.Config
cacheDir string
// mdmSource holds the per-Client MDM policy source and its change
// detector as one unit. Set by SetMDMPolicyFetcher (called from the
// Kotlin side). Each Run passes the loader to the resolved Config so
// applyMDMPolicy picks up the active overlay. Nil means "MDM
// enforcement off for this Client".
mdmSource atomic.Pointer[mdmSource]
// Identifies the running profile for the SSO login hint; see profile_state.go.
cfgPath string
stateChangeMu sync.Mutex
stateChangeSubID string
eventSub *peer.EventSubscription
// Closed to stop the watch goroutines from delivering buffered items to a
// Closed to stop the watch goroutine from delivering buffered ticks to a
// listener that has been removed or replaced. See stopStateChangeWatchLocked.
stateChangeDone chan struct{}
@@ -178,6 +186,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
if err != nil {
return err
}
c.applyMDMOverlay(cfg)
c.recorder.UpdateManagementAddress(cfg.ManagementURL.String())
c.recorder.UpdateRosenpass(cfg.RosenpassEnabled, cfg.RosenpassPermissive)
@@ -203,6 +212,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
internal.WithNetEvents(c.netMgr))
c.setState(cfg, cacheDir, cfgFile, connectClient)
connectClient.SetSyncResponsePersistence(true)
// This path runs the interactive SSO flow, so reaching here means the peer
// is authenticated again — release the latch Status() reports from. Clear
// only once the fresh connect client is installed: until then Status()
@@ -229,6 +239,7 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR
if err != nil {
return err
}
c.applyMDMOverlay(cfg)
c.recorder.UpdateManagementAddress(cfg.ManagementURL.String())
c.recorder.UpdateRosenpass(cfg.RosenpassEnabled, cfg.RosenpassPermissive)
@@ -245,6 +256,7 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
internal.WithNetEvents(c.netMgr))
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)
}
@@ -316,6 +328,19 @@ func (c *Client) NotifyNetworkChange() {
// or "strict"; strict also anonymizes internal IP ranges, peer names, and
// WireGuard public keys, and implies anonymize.
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()
// If the engine hasn't been started, load config from disk
@@ -327,9 +352,15 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
if err != nil {
return "", fmt.Errorf("load config: %w", err)
}
c.applyMDMOverlay(cfg)
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{
InternalConfig: cfg,
StatusRecorder: c.recorder,
@@ -367,6 +398,9 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
if err != nil {
return "", fmt.Errorf("generate debug bundle: %w", err)
}
if !upload {
return debug.ExportBundle(path)
}
defer func() {
if err := os.Remove(path); err != nil {
log.Errorf("failed to remove debug bundle file: %v", err)
@@ -463,6 +497,7 @@ func (c *Client) Networks() *NetworkArray {
routesMap := routeManager.GetClientRoutesWithNetID()
v6Merged := route.V6ExitMergeSet(routesMap)
resolvedDomains := c.recorder.GetResolvedDomainsStates()
activeRoutePeers := c.recorder.GetActiveRoutePeers()
networkArray := &NetworkArray{
items: make([]Network, 0),
@@ -476,7 +511,7 @@ func (c *Client) Networks() *NetworkArray {
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 {
continue
}
@@ -485,14 +520,14 @@ func (c *Client) Networks() *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]
netStr := r.Network.String()
if r.IsDynamic() {
netStr = r.Domains.SafeString()
}
routePeer, err := c.findBestRoutePeer(routes)
routePeer, err := c.findBestRoutePeer(routes, activeRoutePeers)
if err != nil {
log.Errorf("could not get peer info for route %s: %v", id, err)
return nil
@@ -516,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
// HA route group. Falls back to the first connected peer, then the first peer.
func (c *Client) findBestRoutePeer(routes []*route.Route) (peer.State, error) {
netStr := routes[0].Network.String()
fullStatus := c.recorder.GetFullStatus()
for _, p := range fullStatus.Peers {
if _, ok := p.GetRoutes()[netStr]; ok {
func (c *Client) findBestRoutePeer(routes []*route.Route, activeRoutePeers map[route.HAUniqueID]string) (peer.State, error) {
if peerKey, ok := activeRoutePeers[routes[0].GetHAUniqueID()]; ok {
if p, err := c.recorder.GetPeer(peerKey); err == nil {
return p, nil
}
}
+52
View File
@@ -0,0 +1,52 @@
//go:build android
package android
import (
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/mdm"
)
type mdmSource struct {
loader *mdm.Loader
detector *mdm.ChangeDetector
}
// SetMDMPolicyFetcher registers the native-provided MDM policy fetcher on
// this Client; passing nil disables MDM enforcement.
func (c *Client) SetMDMPolicyFetcher(p PolicyFetcher) {
loader := loaderFor(p)
c.mdmSource.Store(&mdmSource{loader: loader, detector: mdm.NewChangeDetector(loader)})
}
// HasMDMPolicyChanged re-reads the managed configuration and reports whether
// it changed since the last observation; call it from the native OS-change
// notification and restart the engine only on true.
func (c *Client) HasMDMPolicyChanged() bool {
src := c.mdmSource.Load()
if src == nil {
return false
}
return src.detector.Changed()
}
// GetRestrictionsJSON returns the UI enforcement snapshot derived from the
// active MDM policy, in the JSON shape shared with the desktop frontend.
func (c *Client) GetRestrictionsJSON() (string, error) {
return mdm.BuildRestrictions(c.mdmLoader().Load()).JSON()
}
func (c *Client) applyMDMOverlay(cfg *profilemanager.Config) {
loader := c.mdmLoader()
if cfg == nil || loader == nil {
return
}
cfg.ApplyMDMPolicy(loader.Load())
}
func (c *Client) mdmLoader() *mdm.Loader {
if src := c.mdmSource.Load(); src != nil {
return src.loader
}
return nil
}
+56 -18
View File
@@ -8,6 +8,7 @@ import (
"github.com/netbirdio/netbird/client/internal/auth"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/mdm"
"github.com/netbirdio/netbird/client/mobile"
"github.com/netbirdio/netbird/client/system"
)
@@ -46,16 +47,24 @@ type Auth struct {
// an earlier call is orphaned on the server. It also breaks a client that enrols and then runs from
// the persisted config, because the identity it registered is not the one it runs with — the
// management stream rejects it with "no peer auth method provided".
func NewAuth(cfgPath string, mgmURL string) (*Auth, error) {
inputCfg := profilemanager.ConfigInput{
ConfigPath: cfgPath,
ManagementURL: mgmURL,
//
// Auth is constructed under the active MDM policy: the policy is overlaid on
// the resolved config so the login runs against the enforced values, while
// the persisted config keeps the caller-supplied ones; a caller-supplied
// management URL is ignored while MDM manages that key. A nil fetcher
// disables MDM enforcement.
func NewAuth(cfgPath string, mgmURL string, fetcher PolicyFetcher) (*Auth, error) {
policy := loaderFor(fetcher).Load()
inputCfg := profilemanager.ConfigInput{ConfigPath: cfgPath}
if _, managed := policy.GetString(mdm.KeyManagementURL); !managed {
inputCfg.ManagementURL = mgmURL
}
cfg, err := profilemanager.UpdateOrCreateConfig(inputCfg)
if err != nil {
return nil, err
}
cfg.ApplyMDMPolicy(policy)
return &Auth{
ctx: context.Background(),
@@ -75,9 +84,7 @@ func NewAuthWithConfig(ctx context.Context, config *profilemanager.Config, cfgPa
}
}
// SaveConfigIfSSOSupported test the connectivity with the management server by retrieving the server device flow info.
// If it returns a flow info than save the configuration and return true. If it gets a codes.NotFound, it means that SSO
// is not supported and returns false without saving the configuration. For other errors return false.
// SaveConfigIfSSOSupported reports whether the management server supports SSO; the config is already persisted by NewAuth.
func (a *Auth) SaveConfigIfSSOSupported(listener SSOListener) {
go func() {
sso, err := a.saveConfigIfSSOSupported()
@@ -101,15 +108,10 @@ func (a *Auth) saveConfigIfSSOSupported() (bool, error) {
return false, fmt.Errorf("failed to check SSO support: %v", err)
}
if !supportsSSO {
return false, nil
}
err = profilemanager.WriteOutConfig(a.cfgPath, a.config)
return true, err
return supportsSSO, nil
}
// LoginWithSetupKeyAndSaveConfig test the connectivity with the management server with the setup key.
// LoginWithSetupKeyAndSaveConfig registers the peer with the setup key; the config is already persisted by NewAuth.
func (a *Auth) LoginWithSetupKeyAndSaveConfig(resultListener ErrListener, setupKey string, deviceName string) {
go func() {
err := a.loginWithSetupKeyAndSaveConfig(setupKey, deviceName)
@@ -134,8 +136,7 @@ func (a *Auth) loginWithSetupKeyAndSaveConfig(setupKey string, deviceName string
if err != nil {
return fmt.Errorf("login failed: %v", err)
}
return profilemanager.WriteOutConfig(a.cfgPath, a.config)
return nil
}
// Login try register the client on the server
@@ -193,12 +194,49 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error {
}
func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, isAndroidTV bool) (*auth.TokenInfo, error) {
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV, profileLoginHint(a.cfgPath))
return a.foregroundGetTokenInfoFlow(authClient, urlOpener, isAndroidTV, false)
}
// foregroundGetTokenInfoFlow runs the interactive flow. sessionExtend tells the
// server the token will renew this peer's session rather than log a peer in, so
// it can rule out a silent authorization the IdP could answer from an unrelated
// account. See PKCEAuthorizationFlowRequest.
func (a *Auth) foregroundGetTokenInfoFlow(authClient *auth.Auth, urlOpener URLOpener, isAndroidTV bool, sessionExtend bool) (*auth.TokenInfo, error) {
hint := profileLoginHint(a.cfgPath)
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV, sessionExtend, hint)
if err != nil {
return nil, fmt.Errorf("failed to get OAuth flow: %v", err)
}
return runOAuthFlow(a.ctx, oAuthFlow, urlOpener, nil)
tokenInfo, err := runOAuthFlow(a.ctx, oAuthFlow, urlOpener, nil)
if err != nil {
return nil, err
}
if tokenInfo.MatchesAccount(hint) {
return tokenInfo, nil
}
// The IdP answered from a session belonging to another account. Retrying is
// what makes this recoverable: on a peer already registered the server would
// reject the token, and on a fresh one it would silently register the peer
// under the wrong account and bind the profile to it.
log.Infof("login returned an account other than the one this profile is bound to, retrying with an account prompt")
retryFlow := auth.RetryFlowForAccount(oAuthFlow)
if retryFlow == nil {
return tokenInfo, nil
}
retryToken, err := runOAuthFlow(a.ctx, retryFlow, urlOpener, nil)
if err != nil {
return nil, err
}
if !retryToken.MatchesAccount(hint) {
log.Warnf("login still returned a different account after the prompt, continuing with it")
}
return retryToken, nil
}
// profileLoginHint returns the stored account email for the profile at cfgPath.
+3 -3
View File
@@ -16,7 +16,7 @@ import (
func TestNewAuth_ReusesPersistedIdentity(t *testing.T) {
cfgPath := filepath.Join(t.TempDir(), "config.json")
first, err := NewAuth(cfgPath, "https://api.example.com:443")
first, err := NewAuth(cfgPath, "https://api.example.com:443", nil)
if err != nil {
t.Fatalf("first NewAuth: %v", err)
}
@@ -24,7 +24,7 @@ func TestNewAuth_ReusesPersistedIdentity(t *testing.T) {
t.Fatal("first NewAuth produced no private key")
}
second, err := NewAuth(cfgPath, "https://api.example.com:443")
second, err := NewAuth(cfgPath, "https://api.example.com:443", nil)
if err != nil {
t.Fatalf("second NewAuth: %v", err)
}
@@ -38,7 +38,7 @@ func TestNewAuth_ReusesPersistedIdentity(t *testing.T) {
func TestNewAuth_CreatesConfigWhenAbsent(t *testing.T) {
cfgPath := filepath.Join(t.TempDir(), "config.json")
auth, err := NewAuth(cfgPath, "https://api.example.com:443")
auth, err := NewAuth(cfgPath, "https://api.example.com:443", nil)
if err != nil {
t.Fatalf("NewAuth: %v", err)
}
+19
View File
@@ -0,0 +1,19 @@
package android
import (
"github.com/netbirdio/netbird/client/mdm"
)
// PolicyFetcher is implemented by the native layer to return the current
// managed configuration as a JSON-encoded object string; "" means no MDM
// source is present.
type PolicyFetcher interface {
FetchJSON() string
}
func loaderFor(p PolicyFetcher) *mdm.Loader {
if p == nil {
return mdm.NewJSONLoader(nil)
}
return mdm.NewJSONLoader(p.FetchJSON)
}
+76 -26
View File
@@ -1,12 +1,16 @@
package android
import (
"sync/atomic"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/mdm"
)
// Preferences exports a subset of the internal config for gomobile
type Preferences struct {
configInput profilemanager.ConfigInput
mdmLoader atomic.Pointer[mdm.Loader]
}
// NewPreferences creates a new Preferences instance
@@ -14,20 +18,39 @@ func NewPreferences(configPath string) *Preferences {
ci := profilemanager.ConfigInput{
ConfigPath: configPath,
}
return &Preferences{ci}
return &Preferences{configInput: ci}
}
// SetMDMPolicyFetcher registers the native-provided MDM policy fetcher on
// this Preferences instance; passing nil disables MDM enforcement.
func (p *Preferences) SetMDMPolicyFetcher(f PolicyFetcher) {
p.mdmLoader.Store(loaderFor(f))
}
// GetRestrictionsJSON returns the UI enforcement snapshot derived from the
// active MDM policy, in the JSON shape shared with the desktop frontend.
func (p *Preferences) GetRestrictionsJSON() (string, error) {
return mdm.BuildRestrictions(p.policy()).JSON()
}
func (p *Preferences) policy() *mdm.Policy {
return p.mdmLoader.Load().Load()
}
// GetManagementURL reads URL from config file
func (p *Preferences) GetManagementURL() (string, error) {
if v, ok := p.policy().GetString(mdm.KeyManagementURL); ok {
return mdm.CanonicalURL(v), nil
}
if p.configInput.ManagementURL != "" {
return p.configInput.ManagementURL, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return "", err
}
return cfg.ManagementURL.String(), err
return cfg.ManagementURL.String(), nil
}
// SetManagementURL stores the given URL and waits for commit
@@ -41,7 +64,7 @@ func (p *Preferences) GetAdminURL() (string, error) {
return p.configInput.AdminURL, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return "", err
}
@@ -53,17 +76,21 @@ func (p *Preferences) SetAdminURL(url string) {
p.configInput.AdminURL = url
}
// GetPreSharedKey reads pre-shared key from config file
func (p *Preferences) GetPreSharedKey() (string, error) {
// HasPreSharedKey reports whether a pre-shared key is staged, persisted, or
// enforced by MDM; the key itself is never handed to the native layer.
func (p *Preferences) HasPreSharedKey() (bool, error) {
if _, ok := p.policy().GetString(mdm.KeyPreSharedKey); ok {
return true, nil
}
if p.configInput.PreSharedKey != nil {
return *p.configInput.PreSharedKey, nil
return *p.configInput.PreSharedKey != "", nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return "", err
return false, err
}
return cfg.PreSharedKey, err
return cfg.PreSharedKey != "", nil
}
// SetPreSharedKey stores the given key and waits for commit
@@ -78,11 +105,14 @@ func (p *Preferences) SetRosenpassEnabled(enabled bool) {
// GetRosenpassEnabled reads Rosenpass enabled status from config file
func (p *Preferences) GetRosenpassEnabled() (bool, error) {
if v, ok := p.policy().GetBool(mdm.KeyRosenpassEnabled); ok {
return v, nil
}
if p.configInput.RosenpassEnabled != nil {
return *p.configInput.RosenpassEnabled, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -96,11 +126,14 @@ func (p *Preferences) SetRosenpassPermissive(permissive bool) {
// GetRosenpassPermissive reads Rosenpass permissive setting from config file
func (p *Preferences) GetRosenpassPermissive() (bool, error) {
if v, ok := p.policy().GetBool(mdm.KeyRosenpassPermissive); ok {
return v, nil
}
if p.configInput.RosenpassPermissive != nil {
return *p.configInput.RosenpassPermissive, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -109,11 +142,14 @@ func (p *Preferences) GetRosenpassPermissive() (bool, error) {
// GetDisableClientRoutes reads disable client routes setting from config file
func (p *Preferences) GetDisableClientRoutes() (bool, error) {
if v, ok := p.policy().GetBool(mdm.KeyDisableClientRoutes); ok {
return v, nil
}
if p.configInput.DisableClientRoutes != nil {
return *p.configInput.DisableClientRoutes, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -127,11 +163,14 @@ func (p *Preferences) SetDisableClientRoutes(disable bool) {
// GetDisableServerRoutes reads disable server routes setting from config file
func (p *Preferences) GetDisableServerRoutes() (bool, error) {
if v, ok := p.policy().GetBool(mdm.KeyDisableServerRoutes); ok {
return v, nil
}
if p.configInput.DisableServerRoutes != nil {
return *p.configInput.DisableServerRoutes, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -149,7 +188,7 @@ func (p *Preferences) GetDisableDNS() (bool, error) {
return *p.configInput.DisableDNS, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -167,7 +206,7 @@ func (p *Preferences) GetDisableFirewall() (bool, error) {
return *p.configInput.DisableFirewall, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -181,11 +220,14 @@ func (p *Preferences) SetDisableFirewall(disable bool) {
// GetServerSSHAllowed reads server SSH allowed setting from config file
func (p *Preferences) GetServerSSHAllowed() (bool, error) {
if v, ok := p.policy().GetBool(mdm.KeyAllowServerSSH); ok {
return v, nil
}
if p.configInput.ServerSSHAllowed != nil {
return *p.configInput.ServerSSHAllowed, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -207,7 +249,7 @@ func (p *Preferences) GetEnableSSHRoot() (bool, error) {
return *p.configInput.EnableSSHRoot, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -229,7 +271,7 @@ func (p *Preferences) GetEnableSSHSFTP() (bool, error) {
return *p.configInput.EnableSSHSFTP, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -251,7 +293,7 @@ func (p *Preferences) GetEnableSSHLocalPortForwarding() (bool, error) {
return *p.configInput.EnableSSHLocalPortForwarding, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -273,7 +315,7 @@ func (p *Preferences) GetEnableSSHRemotePortForwarding() (bool, error) {
return *p.configInput.EnableSSHRemotePortForwarding, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -291,11 +333,14 @@ func (p *Preferences) SetEnableSSHRemotePortForwarding(enabled bool) {
// GetBlockInbound reads block inbound setting from config file
func (p *Preferences) GetBlockInbound() (bool, error) {
if v, ok := p.policy().GetBool(mdm.KeyBlockInbound); ok {
return v, nil
}
if p.configInput.BlockInbound != nil {
return *p.configInput.BlockInbound, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -313,7 +358,7 @@ func (p *Preferences) GetDisableIPv6() (bool, error) {
return *p.configInput.DisableIPv6, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
@@ -327,18 +372,20 @@ func (p *Preferences) SetDisableIPv6(disable bool) {
// GetRemoteJobsAllowed reads the remote jobs opt-in from config file
func (p *Preferences) GetRemoteJobsAllowed() (bool, error) {
if p.configInput.RemoteJobsAllowed != nil {
policy := p.policy()
if !policy.HasKey(mdm.KeyRemoteJobsAllowed) && p.configInput.RemoteJobsAllowed != nil {
return *p.configInput.RemoteJobsAllowed, nil
}
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil {
return false, err
}
cfg.ApplyMDMPolicy(policy)
if cfg.RemoteJobsAllowed == nil {
return false, nil
}
return *cfg.RemoteJobsAllowed, err
return *cfg.RemoteJobsAllowed, nil
}
// SetRemoteJobsAllowed stores the given value and waits for commit
@@ -348,6 +395,9 @@ func (p *Preferences) SetRemoteJobsAllowed(allowed bool) {
// Commit writes out the changes to the config file
func (p *Preferences) Commit() error {
if err := profilemanager.CheckMDMConflicts(p.configInput, p.policy()); err != nil {
return err
}
_, err := profilemanager.UpdateOrCreateConfig(p.configInput)
return err
}
+12 -13
View File
@@ -28,14 +28,13 @@ func TestPreferences_DefaultValues(t *testing.T) {
t.Errorf("invalid default management url: %s", defaultVar)
}
var preSharedKey string
preSharedKey, err = p.GetPreSharedKey()
hasPSK, err := p.HasPreSharedKey()
if err != nil {
t.Fatalf("failed to read default preshared key: %s", err)
t.Fatalf("failed to read default preshared key presence: %s", err)
}
if preSharedKey != "" {
t.Errorf("invalid preshared key: %s", preSharedKey)
if hasPSK {
t.Errorf("unexpected preshared key presence on fresh config")
}
}
@@ -65,13 +64,13 @@ func TestPreferences_ReadUncommitedValues(t *testing.T) {
}
p.SetPreSharedKey(exampleString)
resp, err = p.GetPreSharedKey()
hasPSK, err := p.HasPreSharedKey()
if err != nil {
t.Fatalf("failed to read preshared key: %s", err)
t.Fatalf("failed to read preshared key presence: %s", err)
}
if resp != exampleString {
t.Errorf("unexpected preshared key: %s", resp)
if !hasPSK {
t.Errorf("expected preshared key presence after staging one")
}
}
@@ -109,12 +108,12 @@ func TestPreferences_Commit(t *testing.T) {
t.Errorf("unexpected management url: %s", resp)
}
resp, err = p.GetPreSharedKey()
hasPSK, err := p.HasPreSharedKey()
if err != nil {
t.Fatalf("failed to read preshared key: %s", err)
t.Fatalf("failed to read preshared key presence: %s", err)
}
if resp != examplePresharedKey {
t.Errorf("unexpected preshared key: %s", resp)
if !hasPSK {
t.Errorf("expected preshared key presence after commit")
}
}
+6
View File
@@ -54,6 +54,12 @@ func NewProfileManager(configDir string) *ProfileManager {
return &ProfileManager{impl: mobile.NewProfileManager(configDir, androidUsername)}
}
// SetMDMPolicyFetcher registers the native-provided MDM policy fetcher on
// this ProfileManager; passing nil disables MDM enforcement.
func (pm *ProfileManager) SetMDMPolicyFetcher(f PolicyFetcher) {
pm.impl.SetMDMLoader(loaderFor(f))
}
// ListProfiles returns all available profiles, including the default profile,
// with their active status set.
func (pm *ProfileManager) ListProfiles() (*ProfileArray, error) {
+16 -85
View File
@@ -6,13 +6,8 @@ import (
"context"
"fmt"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal"
"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.
@@ -21,16 +16,11 @@ import (
// changed: connection state, the run-loop status label (e.g. NeedsLogin) or
// the session deadline. It mirrors the daemon's SubscribeStatus stream
// trigger — on each signal the consumer pulls the fresh values via
// Status() / SessionExpiresAtUnix().
//
// OnSessionExpiring forwards the engine's session-expiry warnings, fired at
// 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.
// Status() / SessionExpiresAtUnix(). The engine arms no expiry-warning
// timers on Android; the app schedules the warnings from the deadline it
// reads here.
type StateChangeListener interface {
OnStateChanged()
OnSessionExpiring(expiresAtUnix int64, leadMinutes int64, finalWarning bool)
}
// Status returns the connect run-loop's status label — the same value the
@@ -110,11 +100,11 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) {
return
}
// Both subscriptions are buffered (one pending tick, ten pending events),
// so unsubscribing is not enough to stop callbacks: the loops would drain
// what is already queued and deliver it to a listener the caller has
// already removed or replaced. Gate every callback on this registration's
// own signal, which is closed before unsubscribing.
// The subscription is buffered (one pending tick), so unsubscribing is
// not enough to stop callbacks: the loop would drain what is already
// queued and deliver it to a listener the caller has already removed or
// replaced. Gate every callback on this registration's own signal, which
// is closed before unsubscribing.
done := make(chan struct{})
c.stateChangeDone = done
@@ -133,9 +123,6 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) {
listener.OnStateChanged()
}
}()
c.eventSub = c.recorder.SubscribeToEvents()
go watchSessionWarnings(c.eventSub, listener, done)
}
// RemoveStateChangeListener unregisters the state notification listener.
@@ -145,21 +132,6 @@ func (c *Client) RemoveStateChangeListener() {
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
// asks the management server to extend the session deadline. The tunnel is
// untouched: no resync, no reconnect. Async; the result arrives on the
@@ -201,8 +173,8 @@ func (c *Client) CancelExtendAuthSession() {
}
func (c *Client) stopStateChangeWatchLocked() {
// Signal first, unsubscribe second: closing the channels only stops new
// items, and the loops would still hand whatever is buffered to a listener
// Signal first, unsubscribe second: closing the channel only stops new
// items, and the loop would still hand whatever is buffered to a listener
// that is no longer registered.
if c.stateChangeDone != nil {
close(c.stateChangeDone)
@@ -212,49 +184,6 @@ func (c *Client) stopStateChangeWatchLocked() {
c.recorder.UnsubscribeFromStateChanges(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) {
@@ -293,11 +222,13 @@ func (c *Client) extendAuthSession(ctx context.Context, urlOpener URLOpener, isA
}
defer authClient.Close()
// Passing the config path makes the flow pick up the login_hint: an extend
// renews the session of the account already signed in, so it must not stop to
// offer a choice.
// Passing the config path makes the flow pick up the login_hint. That alone
// cannot keep the IdP on this profile's account though — a hint is only a
// suggestion, and a silent authorization is answered from whatever session the
// IdP already has, which need not be this peer's when several accounts are
// signed in. Marking the flow as an extend lets the server rule that out.
a := NewAuthWithConfig(ctx, cfg, cfgPath)
tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV)
tokenInfo, err := a.foregroundGetTokenInfoFlow(authClient, urlOpener, isAndroidTV, true)
if err != nil {
return fmt.Errorf("interactive sso login failed: %v", err)
}
+3 -1
View File
@@ -31,6 +31,8 @@ const (
// PasswordRequiredMarker tells Java to prompt for a password and retry. It is
// a string because gomobile flattens errors to their message, so a sentinel
// value would not survive the binding.
//
//nolint:gosec // G101 false positive: a sentinel marker, not a credential
const PasswordRequiredMarker = "netbird-ssh-password-required"
// HostKeyUnknownMarker tells Java to show the fingerprint and, on confirmation,
@@ -467,7 +469,7 @@ func (s *SSHClient) requestJWTToken(cfg *profilemanager.Config, cfgPath string)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
defer cancel()
flow, err := auth.NewOAuthFlow(ctx, cfg, false, true, profileLoginHint(cfgPath))
flow, err := auth.NewOAuthFlow(ctx, cfg, false, true, profileLoginHint(cfgPath), false)
if err != nil {
return "", fmt.Errorf("create oauth flow: %w", err)
}
+58 -26
View File
@@ -23,7 +23,10 @@ import (
"github.com/netbirdio/netbird/version"
)
const errCloseConnection = "Failed to close connection: %v"
const (
errCloseConnection = "Failed to close connection: %v"
noUpDownFlag = "no-updown"
)
var (
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)
noUpDown, _ := cmd.Flags().GetBool(noUpDownFlag)
initialLogLevel, err := client.GetLogLevel(cmd.Context(), &proto.GetLogLevelRequest{})
if err != nil {
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 {
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
} else {
@@ -284,34 +288,20 @@ func runForDuration(cmd *cobra.Command, args []string) error {
}
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())
if noUpDown {
enableSyncResponsePersistence(cmd, client)
} else {
needsRestoreUp = !stateWasDown
cmd.Println("netbird down")
needsRestoreUp = restartDaemon(cmd, client, stateWasDown)
}
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
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 {
cpuProfilingStarted = true
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 {
cmd.PrintErrf("Failed to restore service down state: %v\n", status.Convert(err).Message())
} else {
@@ -458,6 +448,47 @@ func setSyncResponsePersistence(cmd *cobra.Command, args []string) error {
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 {
ticker := time.NewTicker(1 * time.Second)
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().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(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")
}
-98
View File
@@ -1,98 +0,0 @@
package cmd
import (
"fmt"
"sort"
"github.com/spf13/cobra"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/proto"
)
var forwardingRulesCmd = &cobra.Command{
Use: "forwarding",
Short: "List forwarding rules",
Long: `Commands to list forwarding rules.`,
}
var forwardingRulesListCmd = &cobra.Command{
Use: "list",
Aliases: []string{"ls"},
Short: "List forwarding rules",
Example: " netbird forwarding list",
Long: "Commands to list forwarding rules.",
RunE: listForwardingRules,
}
func listForwardingRules(cmd *cobra.Command, _ []string) error {
conn, err := getClient(cmd)
if err != nil {
return err
}
defer conn.Close()
client := proto.NewDaemonServiceClient(conn)
resp, err := client.ForwardingRules(cmd.Context(), &proto.EmptyRequest{})
if err != nil {
return fmt.Errorf("failed to list network: %v", status.Convert(err).Message())
}
if len(resp.GetRules()) == 0 {
cmd.Println("No forwarding rules available.")
return nil
}
printForwardingRules(cmd, resp.GetRules())
return nil
}
func printForwardingRules(cmd *cobra.Command, rules []*proto.ForwardingRule) {
cmd.Println("Available forwarding rules:")
// Sort rules by translated address
sort.Slice(rules, func(i, j int) bool {
if rules[i].GetTranslatedAddress() != rules[j].GetTranslatedAddress() {
return rules[i].GetTranslatedAddress() < rules[j].GetTranslatedAddress()
}
if rules[i].GetProtocol() != rules[j].GetProtocol() {
return rules[i].GetProtocol() < rules[j].GetProtocol()
}
return getFirstPort(rules[i].GetDestinationPort()) < getFirstPort(rules[j].GetDestinationPort())
})
var lastIP string
for _, rule := range rules {
dPort := portToString(rule.GetDestinationPort())
tPort := portToString(rule.GetTranslatedPort())
if lastIP != rule.GetTranslatedAddress() {
lastIP = rule.GetTranslatedAddress()
cmd.Printf("\nTranslated peer: %s\n", rule.GetTranslatedHostname())
}
cmd.Printf(" Local %s/%s to %s:%s\n", rule.GetProtocol(), dPort, rule.GetTranslatedAddress(), tPort)
}
}
func getFirstPort(portInfo *proto.PortInfo) int {
switch v := portInfo.PortSelection.(type) {
case *proto.PortInfo_Port:
return int(v.Port)
case *proto.PortInfo_Range_:
return int(v.Range.GetStart())
default:
return 0
}
}
func portToString(translatedPort *proto.PortInfo) string {
switch v := translatedPort.PortSelection.(type) {
case *proto.PortInfo_Port:
return fmt.Sprintf("%d", v.Port)
case *proto.PortInfo_Range_:
return fmt.Sprintf("%d-%d", v.Range.GetStart(), v.Range.GetEnd())
default:
return "No port specified"
}
}
+60 -8
View File
@@ -9,12 +9,11 @@ import (
log "github.com/sirupsen/logrus"
"github.com/spf13/cobra"
"golang.org/x/term"
"google.golang.org/grpc/codes"
gstatus "google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/internal/auth"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/mdm"
nbnet "github.com/netbirdio/netbird/client/net"
"github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/client/server"
@@ -144,10 +143,7 @@ func doDaemonLogin(ctx context.Context, cmd *cobra.Command, providedSetupKey str
err = WithBackOff(func() error {
var backOffErr error
loginResp, backOffErr = client.Login(ctx, &loginRequest)
if s, ok := gstatus.FromError(backOffErr); ok && (s.Code() == codes.InvalidArgument ||
s.Code() == codes.PermissionDenied ||
s.Code() == codes.NotFound ||
s.Code() == codes.Unimplemented) {
if terminalLoginError(backOffErr) {
loginErr = backOffErr
return nil
}
@@ -326,10 +322,33 @@ func doForegroundLogin(ctx context.Context, cmd *cobra.Command, setupKey string,
}
config, err := profilemanager.ReadConfig(configFilePath)
config, err := profilemanager.ReadConfigOrDefault(configFilePath)
if err != nil {
return fmt.Errorf("read config file %s: %v", configFilePath, err)
}
// Reading a config does not provision one: this login is about to dial
// management with the profile's identity, so mint the keys if the profile
// has none yet and put them on disk — a key that stayed in memory would
// come back different on the next run and register a second peer.
//
// Before the MDM overlay below, on purpose: the file must keep the
// profile's own values. The overlay is runtime-only and re-derived on
// every load, so persisting it would turn an enforced management URL or
// pre-shared key into one the user appears to own once the policy is
// withdrawn.
if generated, err := config.EnsureIdentity(); err != nil {
return fmt.Errorf("ensure profile identity: %v", err)
} else if generated {
if err := profilemanager.WriteOutConfig(configFilePath, config); err != nil {
return fmt.Errorf("write out config file %s: %v", configFilePath, err)
}
}
// CLI standalone login: profilemanager no longer auto-applies MDM,
// so layer in the OS-native policy here. Desktop builds construct
// a Loader with no fetcher — the build-tagged loadPlatform reads
// the registry/plist directly.
config.ApplyMDMPolicy(mdm.NewLoader(nil).Load())
// Mirror runInForegroundMode: recover residual state (DNS, firewall,
// ssh config, legacy routing) from a previous unclean shutdown and
@@ -406,11 +425,44 @@ func foregroundGetTokenInfo(ctx context.Context, cmd *cobra.Command, config *pro
hint = profileState.Email
}
oAuthFlow, err := auth.NewOAuthFlow(ctx, config, util.HasGraphicalSession(), false, hint)
oAuthFlow, err := auth.NewOAuthFlow(ctx, config, util.HasGraphicalSession(), false, hint, false)
if err != nil {
return nil, err
}
tokenInfo, err := runInteractiveFlow(cmd, oAuthFlow)
if err != nil {
return nil, err
}
if tokenInfo.MatchesAccount(hint) {
return tokenInfo, nil
}
// The IdP answered from a session belonging to another account. Retrying is
// what makes this recoverable: on a peer already registered the server would
// reject the token, and on a fresh one it would silently register the peer
// under the wrong account and bind the profile to it.
cmd.Println("The login returned a different account than this profile uses. Asking to sign in again.")
retryFlow := auth.RetryFlowForAccount(oAuthFlow)
if retryFlow == nil {
return tokenInfo, nil
}
retryToken, err := runInteractiveFlow(cmd, retryFlow)
if err != nil {
return nil, err
}
if !retryToken.MatchesAccount(hint) {
log.Warnf("login still returned a different account after the prompt, continuing with it")
}
return retryToken, nil
}
// runInteractiveFlow requests the authorization info, shows the URL to the user
// and blocks until the token comes back.
func runInteractiveFlow(cmd *cobra.Command, oAuthFlow auth.OAuthFlow) (*auth.TokenInfo, error) {
flowInfo, err := oAuthFlow.RequestAuthInfo(context.TODO())
if err != nil {
return nil, fmt.Errorf("getting a request OAuth flow info failed: %v", err)
+39 -3
View File
@@ -20,6 +20,8 @@ import (
"github.com/spf13/cobra"
"github.com/spf13/pflag"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
gstatus "google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/anonymize"
daddr "github.com/netbirdio/netbird/client/internal/daemonaddr"
@@ -175,7 +177,6 @@ func init() {
rootCmd.AddCommand(versionCmd)
rootCmd.AddCommand(sshCmd)
rootCmd.AddCommand(networksCMD)
rootCmd.AddCommand(forwardingRulesCmd)
rootCmd.AddCommand(debugCmd)
rootCmd.AddCommand(profileCmd)
rootCmd.AddCommand(exposeCmd)
@@ -183,8 +184,6 @@ func init() {
networksCMD.AddCommand(routesListCmd)
networksCMD.AddCommand(routesSelectCmd, routesDeselectCmd)
forwardingRulesCmd.AddCommand(forwardingRulesListCmd)
debugCmd.AddCommand(debugBundleCmd)
debugCmd.AddCommand(logCmd)
logCmd.AddCommand(logLevelCmd)
@@ -285,6 +284,43 @@ func DialClientGRPCServer(ctx context.Context, addr string) (*grpc.ClientConn, e
return grpc.DialContext(ctx, target, opts...)
}
// terminalLoginError reports whether a Login failure is final, so the backoff
// cycle stops and the caller is told what the daemon said instead of "login
// backoff cycle failed" thirty seconds later. Retrying cannot change any of
// these answers: the request is malformed, the caller is not allowed, the
// target does not exist, a precondition on the daemon refuses it (the
// update-settings kill switch, an MDM-managed field), or the method is not
// implemented.
//
// Both `netbird up` and `netbird login` run Login through the backoff, and
// they each carried their own copy of this list — which is how one of them
// ended up retrying a refusal the other treated as final.
func terminalLoginError(err error) bool {
// A successful Login reaches here with a nil error, and that is not a
// terminal failure. Handled explicitly rather than left to
// gstatus.FromError, which answers (nil, true) for a nil error and leans on
// Status.Code tolerating a nil receiver to come back as codes.OK.
if err == nil {
return false
}
s, ok := gstatus.FromError(err)
if !ok {
return false
}
switch s.Code() {
case codes.InvalidArgument,
codes.PermissionDenied,
codes.NotFound,
codes.FailedPrecondition,
codes.Unimplemented:
return true
default:
return false
}
}
// WithBackOff execute function in backoff cycle.
func WithBackOff(bf func() error) error {
return backoff.RetryNotify(bf, CLIBackOffSettings, func(err error, duration time.Duration) {
+50
View File
@@ -7,6 +7,7 @@ import (
"fmt"
"net/http"
"runtime"
"slices"
"strings"
"sync"
@@ -25,6 +26,30 @@ var serviceCmd = &cobra.Command{
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 (
serviceName 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)
}
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
}
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/internal/daemonaddr"
"github.com/netbirdio/netbird/client/internal/elevate"
"github.com/netbirdio/netbird/util"
)
@@ -43,10 +44,33 @@ func serviceParamsPath() string {
// loadServiceParams reads saved service parameters from disk.
// 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) {
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 os.IsNotExist(err) {
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 not set, saved env vars are used entirely.
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 len(params.ServiceEnvVars) > 0 {
if len(saved) > 0 {
// No explicit env vars: rebuild serviceEnvVars from saved params.
serviceEnvVars = envMapToSlice(params.ServiceEnvVars)
serviceEnvVars = envMapToSlice(saved)
}
return
}
@@ -204,13 +234,13 @@ func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
return
}
if len(params.ServiceEnvVars) == 0 {
if len(saved) == 0 {
return
}
// Merge saved values underneath explicit ones.
merged := make(map[string]string, len(params.ServiceEnvVars)+len(explicit))
maps.Copy(merged, params.ServiceEnvVars)
merged := make(map[string]string, len(saved)+len(explicit))
maps.Copy(merged, saved)
maps.Copy(merged, explicit) // explicit wins on conflict
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.
func envMapToSlice(m map[string]string) []string {
s := make([]string, 0, len(m))
+54
View File
@@ -9,6 +9,7 @@ import (
"go/token"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
@@ -353,6 +354,59 @@ func TestApplyServiceEnvParams_NotChanged(t *testing.T) {
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) {
origServiceEnvVars := serviceEnvVars
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)
}
+3 -4
View File
@@ -6,9 +6,9 @@ import (
"testing"
"time"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel"
"go.uber.org/mock/gomock"
"google.golang.org/grpc"
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator"
@@ -28,7 +28,6 @@ import (
mgmt "github.com/netbirdio/netbird/management/server"
"github.com/netbirdio/netbird/management/server/activity"
"github.com/netbirdio/netbird/management/server/groups"
"github.com/netbirdio/netbird/management/server/integrations/port_forwarding"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/settings"
"github.com/netbirdio/netbird/management/server/store"
@@ -124,9 +123,9 @@ func startManagement(t *testing.T, config *config.Config, testFile string) (*grp
updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := mgmt.NewAccountRequestBuffer(ctx, store)
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersmanager), config, nil)
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", manager.NewEphemeralManager(store, peersmanager), config, nil)
accountManager, err := mgmt.BuildManager(ctx, config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManagerMock, false, cacheStore)
accountManager, err := mgmt.BuildManager(ctx, config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, settingsMockManager, permissionsManagerMock, false, cacheStore)
if err != nil {
t.Fatal(err)
}
+33 -7
View File
@@ -21,6 +21,7 @@ import (
"github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/mdm"
nbnet "github.com/netbirdio/netbird/client/net"
"github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/client/server"
@@ -234,6 +235,10 @@ func runInForegroundMode(ctx context.Context, cmd *cobra.Command, activeProf *pr
if err != nil {
return fmt.Errorf("get config file: %v", err)
}
// CLI foreground path runs without the daemon Server: layer in the
// active MDM policy explicitly so a forced ManagementURL / PSK /
// other managed key actually takes effect on this run.
config.ApplyMDMPolicy(mdm.NewLoader(nil).Load())
_, _ = profilemanager.UpdateOldManagementURL(ctx, config, configFilePath)
@@ -352,9 +357,17 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager
// set the new config
req := setupSetConfigReq(customDNSAddressConverted, cmd, activeProf.ID.String(), username.Username)
if _, err := client.SetConfig(ctx, req); err != nil {
if st, ok := gstatus.FromError(err); ok && st.Code() == codes.Unavailable {
log.Warnf("setConfig method is not available in the daemon: %s", st.Message())
} else {
switch reason, refused := refusedSettingsUpdate(err); {
case refused:
// Failing here is the point: carrying on would connect while
// silently dropping the settings the caller asked for, since
// nothing further down the line applies them.
return fmt.Errorf("the daemon refused the settings update: %s", reason)
case gstatus.Code(err) == codes.Unavailable:
// The daemon cannot serve the method at all, which is what this
// code means; an older daemon without it lands here.
log.Warnf("the daemon did not apply the settings update: %s", gstatus.Convert(err).Message())
default:
return daemonCallError("call service setConfig method", err)
}
}
@@ -395,10 +408,7 @@ func doDaemonUp(ctx context.Context, cmd *cobra.Command, client proto.DaemonServ
err = WithBackOff(func() error {
var backOffErr error
loginResp, backOffErr = client.Login(ctx, loginRequest)
if s, ok := gstatus.FromError(backOffErr); ok && (s.Code() == codes.InvalidArgument ||
s.Code() == codes.PermissionDenied ||
s.Code() == codes.NotFound ||
s.Code() == codes.Unimplemented) {
if terminalLoginError(backOffErr) {
loginErr = backOffErr
return nil
}
@@ -467,6 +477,22 @@ func setSSHSetConfigFields(req *proto.SetConfigRequest, cmd *cobra.Command) {
}
}
// refusedSettingsUpdate reports whether err is the daemon refusing the settings
// a request carried — the update-settings kill switch, or a field an MDM policy
// manages — and returns the reason it gave.
//
// The distinction that matters is against codes.Unavailable, which means the
// daemon cannot serve the call: that one is worth a warning, because an older
// daemon without the method lands there and the rest of `netbird up` still
// works. A refusal is not, because the settings would be silently dropped.
func refusedSettingsUpdate(err error) (string, bool) {
st, ok := gstatus.FromError(err)
if !ok || st.Code() != codes.FailedPrecondition {
return "", false
}
return st.Message(), true
}
func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, profileName, username string) *proto.SetConfigRequest {
var req proto.SetConfigRequest
req.ProfileName = profileName
+85
View File
@@ -0,0 +1,85 @@
package cmd
import (
"errors"
"testing"
"github.com/stretchr/testify/require"
"google.golang.org/grpc/codes"
gstatus "google.golang.org/grpc/status"
)
// A refused settings update has to fail `netbird up`, or a caller that asked
// for a setting the daemon will not apply connects as if it had been applied.
// The daemon being unable to serve the call is the case that stays a warning.
func TestRefusedSettingsUpdate(t *testing.T) {
tests := []struct {
name string
err error
wantRefused bool
}{
{
name: "the kill switch refused the change",
err: gstatus.Errorf(codes.FailedPrecondition, "update settings are disabled, you cannot use this feature without update settings enabled"),
wantRefused: true,
},
{
name: "an MDM policy manages the field",
err: gstatus.Errorf(codes.FailedPrecondition, "fields managed by MDM policy: managementURL"),
wantRefused: true,
},
{
name: "the daemon cannot serve the call",
err: gstatus.Errorf(codes.Unavailable, "connection refused"),
wantRefused: false,
},
{
name: "any other RPC failure",
err: gstatus.Errorf(codes.Internal, "boom"),
wantRefused: false,
},
{
name: "not a status error at all",
err: errors.New("boom"),
wantRefused: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
reason, refused := refusedSettingsUpdate(tt.err)
require.Equal(t, tt.wantRefused, refused)
if tt.wantRefused {
require.Equal(t, gstatus.Convert(tt.err).Message(), reason, "the daemon's reason must reach the caller")
}
})
}
}
// Both `netbird up` and `netbird login` drive Login through the backoff cycle,
// and a final answer has to stop it: retrying a refusal only replaces the
// daemon's reason with "login backoff cycle failed" thirty seconds later.
func TestTerminalLoginError(t *testing.T) {
tests := []struct {
name string
err error
wantTerminal bool
}{
{name: "settings refused by the kill switch", err: gstatus.Errorf(codes.FailedPrecondition, "update settings are disabled"), wantTerminal: true},
{name: "field managed by MDM", err: gstatus.Errorf(codes.FailedPrecondition, "fields managed by MDM policy: managementURL"), wantTerminal: true},
{name: "caller not allowed", err: gstatus.Errorf(codes.PermissionDenied, "nope"), wantTerminal: true},
{name: "malformed request", err: gstatus.Errorf(codes.InvalidArgument, "nope"), wantTerminal: true},
{name: "profile not found", err: gstatus.Errorf(codes.NotFound, "nope"), wantTerminal: true},
{name: "method missing on an older daemon", err: gstatus.Errorf(codes.Unimplemented, "nope"), wantTerminal: true},
{name: "daemon unreachable, worth retrying", err: gstatus.Errorf(codes.Unavailable, "connection refused"), wantTerminal: false},
{name: "transient internal failure", err: gstatus.Errorf(codes.Internal, "boom"), wantTerminal: false},
{name: "not a status error", err: errors.New("boom"), wantTerminal: false},
{name: "no error at all, the login succeeded", err: nil, wantTerminal: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
require.Equal(t, tt.wantTerminal, terminalLoginError(tt.err))
})
}
}
+5
View File
@@ -21,6 +21,7 @@ import (
"github.com/netbirdio/netbird/client/internal/auth"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/mdm"
nbssh "github.com/netbirdio/netbird/client/ssh"
"github.com/netbirdio/netbird/client/system"
"github.com/netbirdio/netbird/shared/management/domain"
@@ -229,6 +230,10 @@ func New(opts Options) (*Client, error) {
if err != nil {
return nil, fmt.Errorf("create config: %w", err)
}
// Embedded path runs without the daemon Server: apply the active
// MDM policy explicitly so a forced ManagementURL / PSK / other
// managed key takes effect on this embedded engine instance.
config.ApplyMDMPolicy(mdm.NewLoader(nil).Load())
if opts.PrivateKey != "" {
config.PrivateKey = opts.PrivateKey
+3 -4
View File
@@ -6,8 +6,8 @@ import (
"testing"
"time"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"google.golang.org/grpc"
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller"
@@ -21,7 +21,6 @@ import (
nbcache "github.com/netbirdio/netbird/management/server/cache"
"github.com/netbirdio/netbird/management/server/groups"
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator"
"github.com/netbirdio/netbird/management/server/integrations/port_forwarding"
"github.com/netbirdio/netbird/management/server/job"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/settings"
@@ -146,8 +145,8 @@ func startManagement(t *testing.T, signalAddr string) string {
updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := mgmt.NewAccountRequestBuffer(context.Background(), testStore)
networkMapController := controller.NewController(context.Background(), testStore, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(testStore, peersManager), cfg, nil)
accountManager, err := mgmt.BuildManager(context.Background(), cfg, testStore, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
networkMapController := controller.NewController(context.Background(), testStore, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", manager.NewEphemeralManager(testStore, peersManager), cfg, nil)
accountManager, err := mgmt.BuildManager(context.Background(), cfg, testStore, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, settingsMockManager, permissionsManager, false, cacheStore)
require.NoError(t, err)
secretsManager, err := nbgrpc.NewTimeBasedAuthSecretsManager(updateManager, cfg.TURNConfig, cfg.Relay, settingsMockManager, groupsManager)
-166
View File
@@ -8,177 +8,11 @@ import (
"strconv"
"strings"
"github.com/hashicorp/go-multierror"
log "github.com/sirupsen/logrus"
nberrors "github.com/netbirdio/netbird/client/errors"
firewall "github.com/netbirdio/netbird/client/firewall/manager"
)
func (r *family) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) {
ruleID := rule.ID()
if _, exists := r.rules[ruleID+dnatSuffix]; exists {
return rule, nil
}
toDestination := rule.TranslatedAddress.String()
switch {
case len(rule.TranslatedPort.Values) == 0:
// no translated port, use original port
case len(rule.TranslatedPort.Values) == 1:
toDestination += fmt.Sprintf(":%d", rule.TranslatedPort.Values[0])
case rule.TranslatedPort.IsRange && len(rule.TranslatedPort.Values) == 2:
// need the "/originalport" suffix to avoid dnat port randomization
toDestination += fmt.Sprintf(":%d-%d/%d", rule.TranslatedPort.Values[0], rule.TranslatedPort.Values[1], rule.DestinationPort.Values[0])
default:
return nil, fmt.Errorf("invalid translated port: %v", rule.TranslatedPort)
}
proto := strings.ToLower(string(rule.Protocol))
rules := make(map[firewall.RuleID]ruleInfo, 3)
// DNAT rule
dnatRule := []string{
"!", "-i", r.wgIface.Name(),
"-p", proto,
"-j", "DNAT",
"--to-destination", toDestination,
}
dnatRule = append(dnatRule, applyPort("--dport", &rule.DestinationPort)...)
rules[ruleID+dnatSuffix] = ruleInfo{
table: tableNat,
chain: chainRTRdr,
rule: dnatRule,
}
// SNAT rule
snatRule := []string{
"-o", r.wgIface.Name(),
"-p", proto,
"-d", rule.TranslatedAddress.String(),
"-j", "MASQUERADE",
}
snatRule = append(snatRule, applyPort("--dport", &rule.TranslatedPort)...)
rules[ruleID+snatSuffix] = ruleInfo{
table: tableNat,
chain: chainRTNAT,
rule: snatRule,
}
// Forward filtering rule, if fwd policy is DROP
forwardRule := []string{
"-o", r.wgIface.Name(),
"-p", proto,
"-d", rule.TranslatedAddress.String(),
"-j", "ACCEPT",
}
forwardRule = append(forwardRule, applyPort("--dport", &rule.TranslatedPort)...)
rules[ruleID+fwdSuffix] = ruleInfo{
table: tableFilter,
chain: chainRTFwdOut,
rule: forwardRule,
}
for key, ruleInfo := range rules {
if err := r.iptablesClient.Append(ruleInfo.table, ruleInfo.chain, ruleInfo.rule...); err != nil {
r.cleanupFailedDNATAdd(rules)
return nil, fmt.Errorf("add rule %s: %w", key, err)
}
r.rules[key] = ruleInfo.rule
}
if err := r.ipFwdState.RequestForwarding(r.v6); err != nil {
r.cleanupFailedDNATAdd(rules)
return nil, fmt.Errorf("enable forwarding: %w", err)
}
r.updateState()
return rule, nil
}
// cleanupFailedDNATAdd removes the bookkeeping written by a partially applied
// AddDNATRule before rolling back the kernel rules, so no entries remain that
// never got a forwarding refcount. rollbackRules re-adds entries it failed to
// remove from the kernel.
func (r *family) cleanupFailedDNATAdd(rules map[firewall.RuleID]ruleInfo) {
for key := range rules {
delete(r.rules, key)
}
if err := r.rollbackRules(rules); err != nil {
log.Errorf("rollback failed: %v", err)
}
}
func (r *family) rollbackRules(rules map[firewall.RuleID]ruleInfo) error {
var merr *multierror.Error
for key, ruleInfo := range rules {
if err := r.iptablesClient.DeleteIfExists(ruleInfo.table, ruleInfo.chain, ruleInfo.rule...); err != nil {
merr = multierror.Append(merr, fmt.Errorf("rollback rule %s: %w", key, err))
// On rollback error, add to rules map for next cleanup
r.rules[key] = ruleInfo.rule
}
}
if merr != nil {
r.updateState()
}
return nberrors.FormatErrorOrNil(merr)
}
func (r *family) DeleteDNATRule(rule firewall.Rule) error {
ruleID := rule.ID()
_, hadDNAT := r.rules[ruleID+dnatSuffix]
_, hadSNAT := r.rules[ruleID+snatSuffix]
_, hadFWD := r.rules[ruleID+fwdSuffix]
if !hadDNAT && !hadSNAT && !hadFWD {
return nil
}
var merr *multierror.Error
if dnatRule, exists := r.rules[ruleID+dnatSuffix]; exists {
if err := r.iptablesClient.Delete(tableNat, chainRTRdr, dnatRule...); err != nil {
merr = multierror.Append(merr, fmt.Errorf("delete DNAT rule: %w", err))
} else {
delete(r.rules, ruleID+dnatSuffix)
}
}
if snatRule, exists := r.rules[ruleID+snatSuffix]; exists {
if err := r.iptablesClient.Delete(tableNat, chainRTNAT, snatRule...); err != nil {
merr = multierror.Append(merr, fmt.Errorf("delete SNAT rule: %w", err))
} else {
delete(r.rules, ruleID+snatSuffix)
}
}
if fwdRule, exists := r.rules[ruleID+fwdSuffix]; exists {
if err := r.iptablesClient.Delete(tableFilter, chainRTFwdOut, fwdRule...); err != nil {
merr = multierror.Append(merr, fmt.Errorf("delete forward rule: %w", err))
} else {
delete(r.rules, ruleID+fwdSuffix)
}
}
// Release the refcount only once all rules are gone from the kernel. On
// partial failure the failed entries stay in r.rules so a retry can remove
// them and release then.
if merr == nil {
r.releaseForwarding()
}
r.updateState()
return nberrors.FormatErrorOrNil(merr)
}
// releaseForwarding drops one IP forwarding reference, logging any error.
func (r *family) releaseForwarding() {
if err := r.ipFwdState.ReleaseForwarding(r.v6); err != nil {
log.Errorf("release IP forwarding: %v", err)
}
}
func (r *family) AddInboundDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
@@ -1,240 +0,0 @@
//go:build privileged
package iptables
import (
"net/netip"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
fw "github.com/netbirdio/netbird/client/firewall/manager"
"github.com/netbirdio/netbird/client/iface"
"github.com/netbirdio/netbird/client/iface/wgaddr"
)
func iptRefcountIfaceV4() *iFaceMock {
return &iFaceMock{
NameFunc: func() string { return "wt-refcount" },
AddressFunc: func() wgaddr.Address {
return wgaddr.Address{
IP: netip.MustParseAddr("10.20.0.1"),
Network: netip.MustParsePrefix("10.20.0.0/24"),
}
},
}
}
func iptRefcountIfaceDual() *iFaceMock {
return &iFaceMock{
NameFunc: func() string { return "wt-refcount" },
AddressFunc: func() wgaddr.Address {
return wgaddr.Address{
IP: netip.MustParseAddr("10.20.0.1"),
Network: netip.MustParsePrefix("10.20.0.0/24"),
IPv6: netip.MustParseAddr("fd00::1"),
IPv6Net: netip.MustParsePrefix("fd00::/64"),
}
},
}
}
func newIptRefcountManager(t *testing.T, dual bool) *Manager {
t.Helper()
var ifMock *iFaceMock
if dual {
ifMock = iptRefcountIfaceDual()
} else {
ifMock = iptRefcountIfaceV4()
}
m, err := Create(ifMock, iface.DefaultMTU)
require.NoError(t, err, "create manager")
require.NoError(t, m.Init(nil), "init manager")
t.Cleanup(func() {
require.NoError(t, m.Close(nil), "close manager")
})
return m
}
func iptDnatV4(port uint16) fw.ForwardRule {
return fw.ForwardRule{
Protocol: fw.ProtocolTCP,
DestinationPort: fw.Port{Values: []uint16{port}},
TranslatedAddress: netip.MustParseAddr("10.20.0.2"),
TranslatedPort: fw.Port{Values: []uint16{80}},
}
}
func iptDnatV6(port uint16) fw.ForwardRule {
return fw.ForwardRule{
Protocol: fw.ProtocolTCP,
DestinationPort: fw.Port{Values: []uint16{port}},
TranslatedAddress: netip.MustParseAddr("fd00::2"),
TranslatedPort: fw.Port{Values: []uint16{80}},
}
}
// TestIptablesRouting_RepeatedEnableSingleReference verifies that EnableRouting
// (called on every network-map update) holds at most one reference per family
// and a single DisableRouting drops both back to zero.
func TestIptablesRouting_RepeatedEnableSingleReference(t *testing.T) {
m := newIptRefcountManager(t, true)
state := m.family4.ipFwdState
require.NoError(t, m.EnableRouting(), "first enable")
require.NoError(t, m.EnableRouting(), "second enable")
require.NoError(t, m.EnableRouting(), "third enable")
v4, v6 := state.Counts()
assert.Equal(t, 1, v4, "repeated enable holds a single v4 reference")
assert.Equal(t, 1, v6, "repeated enable holds a single v6 reference")
require.NoError(t, m.DisableRouting(), "disable")
v4, v6 = state.Counts()
assert.Equal(t, 0, v4, "single disable releases the v4 reference")
assert.Equal(t, 0, v6, "single disable releases the v6 reference")
}
// TestIptablesRouting_DisableKeepsDNATReference verifies that an unpaired
// DisableRouting does not release references held by active DNAT rules.
func TestIptablesRouting_DisableKeepsDNATReference(t *testing.T) {
m := newIptRefcountManager(t, true)
state := m.family4.ipFwdState
r1, err := m.AddDNATRule(iptDnatV6(9095))
require.NoError(t, err, "add v6 dnat")
require.NoError(t, m.DisableRouting(), "unpaired disable")
_, v6 := state.Counts()
assert.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting")
require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat")
_, v6 = state.Counts()
assert.Equal(t, 0, v6, "delete releases the DNAT reference")
}
// TestIptablesDNAT_RefcountBalancedV4 covers a Balanced Add/Delete pair on v4.
func TestIptablesDNAT_RefcountBalancedV4(t *testing.T) {
m := newIptRefcountManager(t, false)
state := m.family4.ipFwdState
r1, err := m.AddDNATRule(iptDnatV4(7081))
require.NoError(t, err, "add v4 dnat 1")
v4, v6 := state.Counts()
assert.Equal(t, 1, v4, "v4 refcount after first add")
assert.Equal(t, 0, v6, "v6 refcount unchanged")
r2, err := m.AddDNATRule(iptDnatV4(7082))
require.NoError(t, err, "add v4 dnat 2")
v4, v6 = state.Counts()
assert.Equal(t, 2, v4, "v4 refcount after second add")
assert.Equal(t, 0, v6, "v6 refcount unchanged")
require.NoError(t, m.DeleteDNATRule(r1))
v4, v6 = state.Counts()
assert.Equal(t, 1, v4, "v4 refcount after first delete")
assert.Equal(t, 0, v6, "v6 refcount unchanged")
require.NoError(t, m.DeleteDNATRule(r2))
v4, v6 = state.Counts()
assert.Equal(t, 0, v4, "v4 refcount after second delete")
assert.Equal(t, 0, v6, "v6 refcount unchanged")
}
// TestIptablesDNAT_RefcountBalancedV6 checks the v6 path increments v6 only and
// decrements back to zero.
func TestIptablesDNAT_RefcountBalancedV6(t *testing.T) {
m := newIptRefcountManager(t, true)
require.NotNil(t, m.family6, "v6 family")
require.Same(t, m.family4.ipFwdState, m.family6.ipFwdState, "shared state")
state := m.family4.ipFwdState
r1, err := m.AddDNATRule(iptDnatV6(9081))
require.NoError(t, err, "add v6 dnat 1")
v4, v6 := state.Counts()
assert.Equal(t, 0, v4)
assert.Equal(t, 1, v6, "v6 refcount after first add")
r2, err := m.AddDNATRule(iptDnatV6(9082))
require.NoError(t, err, "add v6 dnat 2")
v4, v6 = state.Counts()
assert.Equal(t, 0, v4, "v4 refcount unchanged")
assert.Equal(t, 2, v6, "v6 refcount after second add")
require.NoError(t, m.DeleteDNATRule(r1))
v4, v6 = state.Counts()
assert.Equal(t, 0, v4, "v4 refcount unchanged")
assert.Equal(t, 1, v6, "v6 refcount after first delete")
require.NoError(t, m.DeleteDNATRule(r2))
v4, v6 = state.Counts()
assert.Equal(t, 0, v4)
assert.Equal(t, 0, v6, "v6 refcount after second delete")
}
// TestIptablesDNAT_DuplicateAddNoLeak verifies the duplicate-rule path returns
// without bumping the refcount.
func TestIptablesDNAT_DuplicateAddNoLeak(t *testing.T) {
m := newIptRefcountManager(t, true)
state := m.family4.ipFwdState
rule := iptDnatV4(7083)
r1, err := m.AddDNATRule(rule)
require.NoError(t, err)
v4, _ := state.Counts()
assert.Equal(t, 1, v4)
_, err = m.AddDNATRule(rule)
require.NoError(t, err, "duplicate add")
v4, _ = state.Counts()
assert.Equal(t, 1, v4, "duplicate add must not increment")
require.NoError(t, m.DeleteDNATRule(r1))
v4, _ = state.Counts()
assert.Equal(t, 0, v4, "single delete must drop to zero")
}
// TestIptablesDNAT_DeleteMissingNoUnderflow verifies Delete on an unknown rule
// neither errors nor releases the refcount.
func TestIptablesDNAT_DeleteMissingNoUnderflow(t *testing.T) {
m := newIptRefcountManager(t, true)
state := m.family4.ipFwdState
phantom := iptDnatV4(7099)
require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4")
v4, v6 := state.Counts()
assert.Equal(t, 0, v4)
assert.Equal(t, 0, v6)
phantom6 := iptDnatV6(9099)
require.NoError(t, m.DeleteDNATRule(&phantom6), "delete missing v6")
v4, v6 = state.Counts()
assert.Equal(t, 0, v4)
assert.Equal(t, 0, v6)
r1, err := m.AddDNATRule(iptDnatV4(7100))
require.NoError(t, err)
v4, _ = state.Counts()
assert.Equal(t, 1, v4, "real add still increments after phantom delete")
require.NoError(t, m.DeleteDNATRule(r1))
}
// TestIptablesDNAT_DoubleDeleteNoUnderflow verifies a second Delete on the same
// rule is a no-op.
func TestIptablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) {
m := newIptRefcountManager(t, true)
state := m.family4.ipFwdState
r1, err := m.AddDNATRule(iptDnatV6(9083))
require.NoError(t, err)
_, v6 := state.Counts()
assert.Equal(t, 1, v6)
require.NoError(t, m.DeleteDNATRule(r1), "first delete")
_, v6 = state.Counts()
assert.Equal(t, 0, v6)
require.NoError(t, m.DeleteDNATRule(r1), "second delete must be no-op")
_, v6 = state.Counts()
assert.Equal(t, 0, v6, "double delete must not underflow")
}
+2 -4
View File
@@ -24,6 +24,7 @@ const (
tableFilter = "filter"
tableNat = "nat"
tableMangle = "mangle"
tableRaw = "raw"
// chainACLInput is the peer ACL chain that holds installed
// peer-filtering rules.
@@ -34,6 +35,7 @@ const (
mangleForwardKey chainKey = "MANGLE-FORWARD"
chainInput = "INPUT"
chainOutput = "OUTPUT"
chainPostrouting = "POSTROUTING"
chainPrerouting = "PREROUTING"
chainForward = "FORWARD"
@@ -54,10 +56,6 @@ const (
markManglePost = "mark-mangle-post"
matchSet = "--match-set"
dnatSuffix firewall.RuleID = "_dnat"
snatSuffix firewall.RuleID = "_snat"
fwdSuffix firewall.RuleID = "_fwd"
// ipv4TCPHeaderSize is the minimum IPv4 (20) + TCP (20) header size for MSS calculation.
ipv4TCPHeaderSize = 40
// ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation.
-9
View File
@@ -81,15 +81,6 @@ func (r *family) hasRule(id nbid.RuleID) bool {
return ok
}
// hasDNATRule reports whether this family owns the DNAT rule set for
// the given user id. DNAT rules live in r.rules under the well-known
// "<id>_dnat" key; the lookup here is used by Manager.DeleteDNATRule
// to pick the right family.
func (r *family) hasDNATRule(id firewall.RuleID) bool {
_, ok := r.rules[id+dnatSuffix]
return ok
}
// DeleteFilterRule removes a previously installed filter rule. The
// rule's stored chain/table identify where to delete from; source set
// references are recovered from the spec via findSets and dropped
+2 -164
View File
@@ -25,9 +25,8 @@ type Manager struct {
wgIface iFaceMapper
ipv4Client *iptables.IPTables
family4 *family
rawSupported bool
ipv4Client *iptables.IPTables
family4 *family
// IPv6 counterparts, nil when no v6 overlay
ipv6Client *iptables.IPTables
@@ -108,10 +107,6 @@ func (m *Manager) Init(stateManager *statemanager.Manager) error {
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
// interface in firewalld's trusted zone without a corresponding Close.
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
if err := m.cleanupNoTrackChain(); err != nil {
merr = multierror.Append(merr, fmt.Errorf("cleanup notrack chain: %w", err))
}
if m.hasIPv6() {
if err := m.family6.Reset(); err != nil {
merr = multierror.Append(merr, fmt.Errorf("reset v6 family: %w", err))
@@ -332,31 +323,6 @@ func (m *Manager) DisableRouting() error {
return m.family4.ipFwdState.ReleaseRouting()
}
// AddDNATRule adds a DNAT rule
func (m *Manager) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) {
m.mutex.Lock()
defer m.mutex.Unlock()
if rule.TranslatedAddress.Is6() {
if !m.hasIPv6() {
return nil, fmt.Errorf("add DNAT rule: %w", firewall.ErrIPv6NotInitialized)
}
return m.family6.AddDNATRule(rule)
}
return m.family4.AddDNATRule(rule)
}
// DeleteDNATRule deletes a DNAT rule
func (m *Manager) DeleteDNATRule(rule firewall.Rule) error {
m.mutex.Lock()
defer m.mutex.Unlock()
if m.hasIPv6() && !m.family4.hasDNATRule(rule.ID()) {
return m.family6.DeleteDNATRule(rule)
}
return m.family4.DeleteDNATRule(rule)
}
// UpdateSet updates the set with the given prefixes
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
m.mutex.Lock()
@@ -440,134 +406,6 @@ func (m *Manager) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Proto
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 {
return []string{"-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"}
}
@@ -497,16 +497,6 @@ func TestIptablesCloseRemovesAllState(t *testing.T) {
require.NoError(t, manager.AddNatRule(pair), "add nat rule")
require.NoError(t, manager.EnableRouting(), "enable routing")
// A DNAT redirect, which also holds a forwarding reference.
dnat := fw.ForwardRule{
Protocol: fw.ProtocolTCP,
DestinationPort: fw.Port{Values: []uint16{8080}},
TranslatedAddress: netip.MustParseAddr("10.20.0.44"),
TranslatedPort: fw.Port{Values: []uint16{80}},
}
_, err = manager.AddDNATRule(dnat)
require.NoError(t, err, "add dnat rule")
require.NotEqual(t, before, snapshotIptables(t, ipv4Client), "the manager must have installed state")
// Everything above stays in place, so Close is what has to remove it.
-10
View File
@@ -172,12 +172,6 @@ type Manager interface {
DisableRouting() error
// AddDNATRule adds outbound DNAT rule for forwarding external traffic to the NetBird network.
AddDNATRule(ForwardRule) (Rule, error)
// DeleteDNATRule deletes the outbound DNAT rule.
DeleteDNATRule(Rule) error
// UpdateSet updates the set with the given prefixes
UpdateSet(hash Set, prefixes []netip.Prefix) error
@@ -192,10 +186,6 @@ type Manager interface {
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
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.
-27
View File
@@ -1,27 +0,0 @@
package manager
import (
"fmt"
"net/netip"
)
// ForwardRule todo figure out better place to this to avoid circular imports
type ForwardRule struct {
Protocol Protocol
DestinationPort Port
TranslatedAddress netip.Addr
TranslatedPort Port
}
func (r ForwardRule) ID() RuleID {
id := fmt.Sprintf("%s;%s;%s;%s",
r.Protocol,
r.DestinationPort.String(),
r.TranslatedAddress.String(),
r.TranslatedPort.String())
return RuleID(id)
}
func (r ForwardRule) String() string {
return fmt.Sprintf("protocol: %s, destinationPort: %s, translatedAddress: %s, translatedPort: %s", r.Protocol, r.DestinationPort.String(), r.TranslatedAddress.String(), r.TranslatedPort.String())
}
-321
View File
@@ -9,332 +9,11 @@ import (
"github.com/google/nftables"
"github.com/google/nftables/binaryutil"
"github.com/google/nftables/expr"
"github.com/google/nftables/xt"
"github.com/hashicorp/go-multierror"
log "github.com/sirupsen/logrus"
nberrors "github.com/netbirdio/netbird/client/errors"
firewall "github.com/netbirdio/netbird/client/firewall/manager"
)
func (r *family) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) {
ruleID := rule.ID()
if _, exists := r.rules[ruleID+dnatSuffix]; exists {
return rule, nil
}
protoNum, err := r.af.protoNum(rule.Protocol)
if err != nil {
return nil, fmt.Errorf("convert protocol to number: %w", err)
}
// Request forwarding before queueing rules: addDnatRedirect/addDnatMasq
// buffer netlink messages on r.conn that the next caller's Flush would
// commit if we returned without flushing them ourselves.
if err := r.ipFwdState.RequestForwarding(r.isV6()); err != nil {
return nil, fmt.Errorf("enable forwarding: %w", err)
}
if err := r.addDnatRedirect(rule, protoNum, ruleID); err != nil {
r.releaseForwarding()
return nil, err
}
if err := r.addDnatMasq(rule, protoNum, ruleID); err != nil {
r.releaseForwarding()
delete(r.rules, ruleID+dnatSuffix)
return nil, err
}
// Unlike iptables, there's no point in adding "out" rules in the forward chain here as our policy is ACCEPT.
// To overcome DROP policies in other chains, we'd have to add rules to the chains there.
// We also cannot just add "oif <iface> accept" there and filter in our own table as we don't know what is supposed to be allowed.
// TODO: find chains with drop policies and add rules there
if err := r.conn.Flush(); err != nil {
r.releaseForwarding()
delete(r.rules, ruleID+dnatSuffix)
delete(r.rules, ruleID+snatSuffix)
return nil, fmt.Errorf("flush rules: %w", err)
}
return &rule, nil
}
func (r *family) addDnatRedirect(rule firewall.ForwardRule, protoNum uint8, ruleID firewall.RuleID) error {
dnatExprs := []expr.Any{
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
&expr.Cmp{
Op: expr.CmpOpNeq,
Register: 1,
Data: ifname(r.wgIface.Name()),
},
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{
Op: expr.CmpOpEq,
Register: 1,
Data: []byte{protoNum},
},
&expr.Payload{
DestRegister: 1,
Base: expr.PayloadBaseTransportHeader,
Offset: 2,
Len: 2,
},
}
portExprs, err := r.applyPort(&rule.DestinationPort, false)
if err != nil {
return fmt.Errorf("apply destination port: %w", err)
}
dnatExprs = append(dnatExprs, portExprs...)
// shifted translated port is not supported in nftables, so we hand this over to xtables
if rule.TranslatedPort.IsRange && len(rule.TranslatedPort.Values) == 2 {
if rule.TranslatedPort.Values[0] != rule.DestinationPort.Values[0] ||
rule.TranslatedPort.Values[1] != rule.DestinationPort.Values[1] {
return r.addXTablesRedirect(dnatExprs, ruleID, rule)
}
}
additionalExprs, regProtoMin, regProtoMax, err := r.handleTranslatedPort(rule)
if err != nil {
return err
}
dnatExprs = append(dnatExprs, additionalExprs...)
dnatExprs = append(dnatExprs,
&expr.NAT{
Type: expr.NATTypeDestNAT,
Family: uint32(r.af.tableFamily),
RegAddrMin: 1,
RegProtoMin: regProtoMin,
RegProtoMax: regProtoMax,
},
)
dnatRule := &nftables.Rule{
Table: r.workTable,
Chain: r.chains[chainNameRoutingRdr],
Exprs: dnatExprs,
UserData: []byte(ruleID + dnatSuffix),
}
r.conn.AddRule(dnatRule)
r.rules[ruleID+dnatSuffix] = dnatRule
return nil
}
func (r *family) handleTranslatedPort(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
switch {
case rule.TranslatedPort.IsRange && len(rule.TranslatedPort.Values) == 2:
return r.handlePortRange(rule)
case len(rule.TranslatedPort.Values) == 0:
return r.handleAddressOnly(rule)
case len(rule.TranslatedPort.Values) == 1:
return r.handleSinglePort(rule)
default:
return nil, 0, 0, fmt.Errorf("invalid translated port: %v", rule.TranslatedPort)
}
}
func (r *family) handlePortRange(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
exprs := []expr.Any{
&expr.Immediate{
Register: 1,
Data: rule.TranslatedAddress.AsSlice(),
},
&expr.Immediate{
Register: 2,
Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[0]),
},
&expr.Immediate{
Register: 3,
Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[1]),
},
}
return exprs, 2, 3, nil
}
func (r *family) handleAddressOnly(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
exprs := []expr.Any{
&expr.Immediate{
Register: 1,
Data: rule.TranslatedAddress.AsSlice(),
},
}
return exprs, 0, 0, nil
}
func (r *family) handleSinglePort(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
exprs := []expr.Any{
&expr.Immediate{
Register: 1,
Data: rule.TranslatedAddress.AsSlice(),
},
&expr.Immediate{
Register: 2,
Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[0]),
},
}
return exprs, 2, 0, nil
}
func (r *family) addXTablesRedirect(dnatExprs []expr.Any, ruleID firewall.RuleID, rule firewall.ForwardRule) error {
dnatExprs = append(dnatExprs,
&expr.Counter{},
&expr.Target{
Name: "DNAT",
Rev: 2,
Info: &xt.NatRange2{
NatRange: xt.NatRange{
Flags: uint(xt.NatRangeMapIPs | xt.NatRangeProtoSpecified | xt.NatRangeProtoOffset),
MinIP: rule.TranslatedAddress.AsSlice(),
MaxIP: rule.TranslatedAddress.AsSlice(),
MinPort: rule.TranslatedPort.Values[0],
MaxPort: rule.TranslatedPort.Values[1],
},
BasePort: rule.DestinationPort.Values[0],
},
},
)
natTable := &nftables.Table{
Name: tableNat,
Family: r.af.tableFamily,
}
dnatRule := &nftables.Rule{
Table: natTable,
Chain: &nftables.Chain{
Name: chainNameNatPrerouting,
Table: natTable,
Type: nftables.ChainTypeNAT,
Hooknum: nftables.ChainHookPrerouting,
Priority: nftables.ChainPriorityNATDest,
},
Exprs: dnatExprs,
UserData: []byte(ruleID + dnatSuffix),
}
r.conn.AddRule(dnatRule)
r.rules[ruleID+dnatSuffix] = dnatRule
return nil
}
func (r *family) addDnatMasq(rule firewall.ForwardRule, protoNum uint8, ruleID firewall.RuleID) error {
portExprs, err := r.applyPort(&rule.TranslatedPort, false)
if err != nil {
return fmt.Errorf("apply translated port: %w", err)
}
masqExprs := []expr.Any{
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
&expr.Cmp{
Op: expr.CmpOpEq,
Register: 1,
Data: ifname(r.wgIface.Name()),
},
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{
Op: expr.CmpOpEq,
Register: 1,
Data: []byte{protoNum},
},
&expr.Payload{
DestRegister: 1,
Base: expr.PayloadBaseNetworkHeader,
Offset: r.af.dstAddrOffset,
Len: r.af.addrLen,
},
&expr.Cmp{
Op: expr.CmpOpEq,
Register: 1,
Data: rule.TranslatedAddress.AsSlice(),
},
}
masqExprs = append(masqExprs, portExprs...)
masqExprs = append(masqExprs, &expr.Masq{})
masqRule := &nftables.Rule{
Table: r.workTable,
Chain: r.chains[chainNameRoutingNat],
Exprs: masqExprs,
UserData: []byte(ruleID + snatSuffix),
}
r.conn.AddRule(masqRule)
r.rules[ruleID+snatSuffix] = masqRule
return nil
}
func (r *family) DeleteDNATRule(rule firewall.Rule) error {
ruleID := rule.ID()
if err := r.refreshRulesMap(); err != nil {
return fmt.Errorf(refreshRulesMapError, err)
}
var merr *multierror.Error
var needsFlush bool
var found bool
if dnatRule, exists := r.rules[ruleID+dnatSuffix]; exists {
found = true
if dnatRule.Handle == 0 {
log.Warnf("dnat rule %s has no handle, removing stale entry", ruleID+dnatSuffix)
delete(r.rules, ruleID+dnatSuffix)
} else if err := r.conn.DelRule(dnatRule); err != nil {
merr = multierror.Append(merr, fmt.Errorf("delete dnat rule: %w", err))
} else {
needsFlush = true
}
}
if masqRule, exists := r.rules[ruleID+snatSuffix]; exists {
found = true
if masqRule.Handle == 0 {
log.Warnf("snat rule %s has no handle, removing stale entry", ruleID+snatSuffix)
delete(r.rules, ruleID+snatSuffix)
} else if err := r.conn.DelRule(masqRule); err != nil {
merr = multierror.Append(merr, fmt.Errorf("delete snat rule: %w", err))
} else {
needsFlush = true
}
}
if needsFlush {
if err := r.conn.Flush(); err != nil {
merr = multierror.Append(merr, fmt.Errorf(flushError, err))
}
}
if merr != nil {
return nberrors.FormatErrorOrNil(merr)
}
delete(r.rules, ruleID+dnatSuffix)
delete(r.rules, ruleID+snatSuffix)
// Release once, only if the rule was present and removed.
if found {
r.releaseForwarding()
}
return nil
}
// releaseForwarding drops one IP forwarding reference, logging any error.
func (r *family) releaseForwarding() {
if err := r.ipFwdState.ReleaseForwarding(r.isV6()); err != nil {
log.Errorf("release IP forwarding: %v", err)
}
}
// isV6 reports whether this family handles the IPv6 table.
func (r *family) isV6() bool {
return r.af.tableFamily == nftables.TableFamilyIPv6
}
func (r *family) AddInboundDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
@@ -1,249 +0,0 @@
//go:build privileged
package nftables
import (
"net/netip"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
fw "github.com/netbirdio/netbird/client/firewall/manager"
"github.com/netbirdio/netbird/client/iface"
"github.com/netbirdio/netbird/client/iface/wgaddr"
)
func nftRefcountIfaceV4() *iFaceMock {
return &iFaceMock{
NameFunc: func() string { return "wt-refcount" },
AddressFunc: func() wgaddr.Address {
return wgaddr.Address{
IP: netip.MustParseAddr("100.96.0.1"),
Network: netip.MustParsePrefix("100.96.0.0/16"),
}
},
}
}
func nftRefcountIfaceDual() *iFaceMock {
return &iFaceMock{
NameFunc: func() string { return "wt-refcount" },
AddressFunc: func() wgaddr.Address {
return wgaddr.Address{
IP: netip.MustParseAddr("100.96.0.1"),
Network: netip.MustParsePrefix("100.96.0.0/16"),
IPv6: netip.MustParseAddr("fd00::1"),
IPv6Net: netip.MustParsePrefix("fd00::/64"),
}
},
}
}
func newNftRefcountManager(t *testing.T, dual bool) *Manager {
t.Helper()
if check() != NFTABLES {
t.Skip("nftables not supported on this system")
}
var ifMock *iFaceMock
if dual {
ifMock = nftRefcountIfaceDual()
} else {
ifMock = nftRefcountIfaceV4()
}
m, err := Create(ifMock, iface.DefaultMTU)
require.NoError(t, err, "create manager")
require.NoError(t, m.Init(nil), "init manager")
t.Cleanup(func() {
require.NoError(t, m.Close(nil), "close manager")
})
return m
}
func dnatV4(port uint16) fw.ForwardRule {
return fw.ForwardRule{
Protocol: fw.ProtocolTCP,
DestinationPort: fw.Port{Values: []uint16{port}},
TranslatedAddress: netip.MustParseAddr("100.96.0.2"),
TranslatedPort: fw.Port{Values: []uint16{80}},
}
}
func dnatV6(port uint16) fw.ForwardRule {
return fw.ForwardRule{
Protocol: fw.ProtocolTCP,
DestinationPort: fw.Port{Values: []uint16{port}},
TranslatedAddress: netip.MustParseAddr("fd00::2"),
TranslatedPort: fw.Port{Values: []uint16{80}},
}
}
// TestNftablesDNAT_RefcountBalancedV4 verifies that Add/Delete pairs leave the
// v4 refcount at zero.
func TestNftablesDNAT_RefcountBalancedV4(t *testing.T) {
m := newNftRefcountManager(t, false)
state := m.family4.ipFwdState
r1, err := m.AddDNATRule(dnatV4(8081))
require.NoError(t, err, "add v4 dnat 1")
v4, v6 := state.Counts()
assert.Equal(t, 1, v4, "v4 refcount after first add")
assert.Equal(t, 0, v6, "v6 refcount unchanged")
r2, err := m.AddDNATRule(dnatV4(8082))
require.NoError(t, err, "add v4 dnat 2")
v4, v6 = state.Counts()
assert.Equal(t, 2, v4, "v4 refcount after second add")
assert.Equal(t, 0, v6, "v6 refcount unchanged")
require.NoError(t, m.DeleteDNATRule(r1), "delete v4 dnat 1")
v4, v6 = state.Counts()
assert.Equal(t, 1, v4, "v4 refcount after first delete")
assert.Equal(t, 0, v6, "v6 refcount unchanged")
require.NoError(t, m.DeleteDNATRule(r2), "delete v4 dnat 2")
v4, v6 = state.Counts()
assert.Equal(t, 0, v4, "v4 refcount after second delete")
assert.Equal(t, 0, v6, "v6 refcount unchanged")
}
// TestNftablesDNAT_RefcountBalancedV6 verifies the v6 path increments v6 only
// and decrements back to zero on Delete.
func TestNftablesDNAT_RefcountBalancedV6(t *testing.T) {
m := newNftRefcountManager(t, true)
require.NotNil(t, m.family6, "v6 family")
require.Same(t, m.family4.ipFwdState, m.family6.ipFwdState, "shared state")
state := m.family4.ipFwdState
r1, err := m.AddDNATRule(dnatV6(9091))
require.NoError(t, err, "add v6 dnat 1")
v4, v6 := state.Counts()
assert.Equal(t, 0, v4, "v4 refcount unchanged")
assert.Equal(t, 1, v6, "v6 refcount after first add")
r2, err := m.AddDNATRule(dnatV6(9092))
require.NoError(t, err, "add v6 dnat 2")
v4, v6 = state.Counts()
assert.Equal(t, 0, v4)
assert.Equal(t, 2, v6, "v6 refcount after second add")
require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat 1")
v4, v6 = state.Counts()
assert.Equal(t, 0, v4, "v4 refcount unchanged")
assert.Equal(t, 1, v6, "v6 refcount after first delete")
require.NoError(t, m.DeleteDNATRule(r2), "delete v6 dnat 2")
v4, v6 = state.Counts()
assert.Equal(t, 0, v4)
assert.Equal(t, 0, v6, "v6 refcount after second delete")
}
// TestNftablesDNAT_DuplicateAddNoLeak verifies that a duplicate Add (same
// ForwardRule) does not double-increment the refcount.
func TestNftablesDNAT_DuplicateAddNoLeak(t *testing.T) {
m := newNftRefcountManager(t, true)
state := m.family4.ipFwdState
rule := dnatV4(8083)
r1, err := m.AddDNATRule(rule)
require.NoError(t, err, "add v4 dnat")
v4, _ := state.Counts()
assert.Equal(t, 1, v4)
// duplicate add: same rule ID, must be a no-op for the refcount.
_, err = m.AddDNATRule(rule)
require.NoError(t, err, "duplicate add")
v4, _ = state.Counts()
assert.Equal(t, 1, v4, "duplicate add must not increment")
require.NoError(t, m.DeleteDNATRule(r1), "delete v4 dnat")
v4, _ = state.Counts()
assert.Equal(t, 0, v4, "single delete must drop to zero")
}
// TestNftablesDNAT_DeleteMissingNoUnderflow verifies deleting a rule that was
// never added does not underflow the refcount.
func TestNftablesDNAT_DeleteMissingNoUnderflow(t *testing.T) {
m := newNftRefcountManager(t, true)
state := m.family4.ipFwdState
// Construct a Rule reference for something never added. The router stores
// rules by ID(), and DeleteDNATRule looks them up in r.rules; a missing
// entry must be a no-op rather than calling Release.
phantom := dnatV4(8099)
require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4 dnat")
v4, v6 := state.Counts()
assert.Equal(t, 0, v4, "v4 refcount unaffected by missing delete")
assert.Equal(t, 0, v6, "v6 refcount unaffected")
phantom6 := dnatV6(9099)
require.NoError(t, m.DeleteDNATRule(&phantom6), "delete missing v6 dnat")
v4, v6 = state.Counts()
assert.Equal(t, 0, v4)
assert.Equal(t, 0, v6, "v6 refcount unaffected by missing delete")
// And after a phantom delete, a real add still results in count=1.
r1, err := m.AddDNATRule(dnatV4(8100))
require.NoError(t, err, "add v4 dnat after phantom delete")
v4, _ = state.Counts()
assert.Equal(t, 1, v4, "real add still increments after phantom delete")
require.NoError(t, m.DeleteDNATRule(r1))
}
// TestNftablesRouting_RepeatedEnableSingleReference verifies that EnableRouting
// (called on every network-map update) holds at most one reference per family
// and a single DisableRouting drops both back to zero.
func TestNftablesRouting_RepeatedEnableSingleReference(t *testing.T) {
m := newNftRefcountManager(t, true)
state := m.family4.ipFwdState
require.NoError(t, m.EnableRouting(), "first enable")
require.NoError(t, m.EnableRouting(), "second enable")
require.NoError(t, m.EnableRouting(), "third enable")
v4, v6 := state.Counts()
assert.Equal(t, 1, v4, "repeated enable holds a single v4 reference")
assert.Equal(t, 1, v6, "repeated enable holds a single v6 reference")
require.NoError(t, m.DisableRouting(), "disable")
v4, v6 = state.Counts()
assert.Equal(t, 0, v4, "single disable releases the v4 reference")
assert.Equal(t, 0, v6, "single disable releases the v6 reference")
}
// TestNftablesRouting_DisableKeepsDNATReference verifies that an unpaired
// DisableRouting does not release references held by active DNAT rules.
func TestNftablesRouting_DisableKeepsDNATReference(t *testing.T) {
m := newNftRefcountManager(t, true)
state := m.family4.ipFwdState
r1, err := m.AddDNATRule(dnatV6(9095))
require.NoError(t, err, "add v6 dnat")
require.NoError(t, m.DisableRouting(), "unpaired disable")
_, v6 := state.Counts()
assert.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting")
require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat")
_, v6 = state.Counts()
assert.Equal(t, 0, v6, "delete releases the DNAT reference")
}
// TestNftablesDNAT_DoubleDeleteNoUnderflow verifies that deleting the same rule
// twice does not underflow the refcount (the second delete is a no-op).
func TestNftablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) {
m := newNftRefcountManager(t, true)
state := m.family4.ipFwdState
r1, err := m.AddDNATRule(dnatV6(9093))
require.NoError(t, err)
_, v6 := state.Counts()
assert.Equal(t, 1, v6)
require.NoError(t, m.DeleteDNATRule(r1), "first delete")
_, v6 = state.Counts()
assert.Equal(t, 0, v6)
require.NoError(t, m.DeleteDNATRule(r1), "second delete must be no-op")
_, v6 = state.Counts()
assert.Equal(t, 0, v6, "double delete must not underflow")
}
-8
View File
@@ -24,7 +24,6 @@ const (
tableRaw = "raw"
tableSecurity = "security"
chainNameNatPrerouting = "PREROUTING"
chainNameRoutingFw = "netbird-rt-fwd"
chainNameRoutingNat = "netbird-rt-postrouting"
chainNameRoutingRdr = "netbird-rt-redirect"
@@ -47,9 +46,6 @@ const (
userDataAcceptForwardRuleOif = "frwacceptoif"
userDataAcceptInputRule = "inputaccept"
dnatSuffix firewall.RuleID = "_dnat"
snatSuffix firewall.RuleID = "_snat"
// ipv4TCPHeaderSize is the minimum IPv4 (20) + TCP (20) header size for MSS calculation.
ipv4TCPHeaderSize = 40
// ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation.
@@ -167,10 +163,6 @@ func (r *family) Reset() error {
merr = multierror.Append(merr, err)
}
if err := r.removeNatPreroutingRules(); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove filter prerouting rules: %w", err))
}
return nberrors.FormatErrorOrNil(merr)
}
-5
View File
@@ -197,11 +197,6 @@ func (r *family) hasRule(id firewall.RuleID) bool {
return ok
}
func (r *family) hasDNATRule(id firewall.RuleID) bool {
_, ok := r.rules[id+dnatSuffix]
return ok
}
// DeleteFilterRule removes a previously installed filter rule. Source
// set references are recovered from the stored rule's expressions via
// findSets and dropped from the shared refcounter.
+3 -226
View File
@@ -12,7 +12,6 @@ import (
"github.com/google/nftables/expr"
"github.com/hashicorp/go-multierror"
log "github.com/sirupsen/logrus"
"golang.org/x/sys/unix"
nberrors "github.com/netbirdio/netbird/client/errors"
firewall "github.com/netbirdio/netbird/client/firewall/manager"
@@ -55,9 +54,6 @@ type Manager struct {
// IPv6 counterpart, nil when no v6 overlay.
family6 *family
notrackOutputChain *nftables.Chain
notrackPreroutingChain *nftables.Chain
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
}
@@ -260,7 +252,7 @@ func (m *Manager) DeleteFilterRule(rule firewall.Rule) error {
m.mutex.Lock()
defer m.mutex.Unlock()
fam, err := m.familyForRuleID(rule.ID(), (*family).hasRule, false)
fam, err := m.familyForRuleID(rule.ID(), (*family).hasRule)
if err != nil {
return err
}
@@ -268,11 +260,8 @@ func (m *Manager) DeleteFilterRule(rule firewall.Rule) error {
}
// familyForRuleID picks the family holding the rule with the given id, using
// the supplied lookup. With refresh set, a miss in both cached maps reloads
// the NAT/DNAT rule maps from the kernel once and re-checks before falling
// back to the v4 family. Filter rules are tracked only in memory and have no
// kernel-backed reload, so their callers pass refresh as false.
func (m *Manager) familyForRuleID(id firewall.RuleID, has func(*family, firewall.RuleID) bool, refresh bool) (*family, error) {
// the supplied lookup, and falls back to the v4 family on a miss.
func (m *Manager) familyForRuleID(id firewall.RuleID, has func(*family, firewall.RuleID) bool) (*family, error) {
if has(m.family4, id) {
return m.family4, nil
}
@@ -282,18 +271,6 @@ func (m *Manager) familyForRuleID(id firewall.RuleID, has func(*family, firewall
if has(m.family6, id) {
return m.family6, nil
}
if !refresh {
return m.family4, nil
}
if err := m.family4.refreshRulesMap(); err != nil {
return nil, fmt.Errorf("refresh v4 rules: %w", err)
}
if err := m.family6.refreshRulesMap(); err != nil {
return nil, fmt.Errorf("refresh v6 rules: %w", err)
}
if has(m.family6, id) && !has(m.family4, id) {
return m.family6, nil
}
return m.family4, nil
}
@@ -455,39 +432,9 @@ func (m *Manager) Flush() error {
}
}
if err := m.refreshNoTrackChains(); err != nil {
log.Errorf("failed to refresh notrack chains: %v", err)
}
return nil
}
// AddDNATRule adds a DNAT rule
func (m *Manager) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) {
m.mutex.Lock()
defer m.mutex.Unlock()
if rule.TranslatedAddress.Is6() {
if !m.hasIPv6() {
return nil, fmt.Errorf("add DNAT rule: %w", firewall.ErrIPv6NotInitialized)
}
return m.family6.AddDNATRule(rule)
}
return m.family4.AddDNATRule(rule)
}
// DeleteDNATRule deletes a DNAT rule
func (m *Manager) DeleteDNATRule(rule firewall.Rule) error {
m.mutex.Lock()
defer m.mutex.Unlock()
r, err := m.familyForRuleID(rule.ID(), (*family).hasDNATRule, true)
if err != nil {
return err
}
return r.DeleteDNATRule(rule)
}
// UpdateSet updates the set with the given prefixes
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
m.mutex.Lock()
@@ -571,176 +518,6 @@ func (m *Manager) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Proto
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) {
return m.createWorkTableFamily(nftables.TableFamilyIPv4)
}
@@ -378,18 +378,6 @@ func TestNftablesManagerCompatibilityWithIptables(t *testing.T) {
err = manager.AddNatRule(pair)
require.NoError(t, err, "failed to add NAT rule")
dnatRule, err := manager.AddDNATRule(fw.ForwardRule{
Protocol: fw.ProtocolTCP,
DestinationPort: fw.Port{Values: []uint16{8080}},
TranslatedAddress: netip.MustParseAddr("100.96.0.2"),
TranslatedPort: fw.Port{Values: []uint16{80}},
})
require.NoError(t, err, "failed to add DNAT rule")
t.Cleanup(func() {
require.NoError(t, manager.DeleteDNATRule(dnatRule), "failed to delete DNAT rule")
})
stdout, stderr = runIptablesSave(t)
verifyIptablesOutput(t, stdout, stderr)
}
@@ -453,18 +441,6 @@ func TestNftablesManagerIPv6CompatibilityWithIp6tables(t *testing.T) {
})
require.NoError(t, err, "add v6 NAT rule")
dnatRule, err := manager.AddDNATRule(fw.ForwardRule{
Protocol: fw.ProtocolTCP,
DestinationPort: fw.Port{Values: []uint16{8080}},
TranslatedAddress: netip.MustParseAddr("fd00::2"),
TranslatedPort: fw.Port{Values: []uint16{80}},
})
require.NoError(t, err, "add v6 DNAT rule")
t.Cleanup(func() {
require.NoError(t, manager.DeleteDNATRule(dnatRule), "delete v6 DNAT rule")
})
stdout, stderr := runIptablesSave(t)
verifyIptablesOutput(t, stdout, stderr)
+1 -36
View File
@@ -192,7 +192,7 @@ func (r *family) addPostroutingRules() {
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{
Key: expr.MetaKeyOIFNAME,
Register: 1,
@@ -459,41 +459,6 @@ func (r *family) RemoveAllLegacyRouteRules() error {
return nberrors.FormatErrorOrNil(merr)
}
func (r *family) removeNatPreroutingRules() error {
table := &nftables.Table{
Name: tableNat,
Family: r.af.tableFamily,
}
chain := &nftables.Chain{
Name: chainNameNatPrerouting,
Table: table,
Hooknum: nftables.ChainHookPrerouting,
Priority: nftables.ChainPriorityNATDest,
Type: nftables.ChainTypeNAT,
}
rules, err := r.conn.GetRules(table, chain)
if err != nil {
return fmt.Errorf("get rules from nat table: %w", err)
}
var merr *multierror.Error
// Delete rules that have our UserData suffix
for _, rule := range rules {
if len(rule.UserData) == 0 || !strings.HasSuffix(string(rule.UserData), string(dnatSuffix)) {
continue
}
if err := r.conn.DelRule(rule); err != nil {
merr = multierror.Append(merr, fmt.Errorf("delete rule %s: %w", rule.UserData, err))
}
}
if err := r.conn.Flush(); err != nil {
merr = multierror.Append(merr, fmt.Errorf(flushError, err))
}
return nberrors.FormatErrorOrNil(merr)
}
func (r *family) RemoveNatRule(pair firewall.RouterPair) error {
if err := r.refreshRulesMap(); err != nil {
return fmt.Errorf(refreshRulesMapError, err)
-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
// by merging the existing prefixes with the new ones, then deduplicating.
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
@@ -9,6 +9,7 @@ import (
log "github.com/sirupsen/logrus"
nberrors "github.com/netbirdio/netbird/client/errors"
"github.com/netbirdio/netbird/client/internal/wincmd"
)
type action string
@@ -91,7 +92,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err
if action == addRule {
args = append(args, extraArgs...)
}
netshCmd := GetSystem32Command("netsh")
netshCmd := wincmd.System32("netsh")
cmd := exec.Command(netshCmd, args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
return cmd.Run()
@@ -100,7 +101,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err
func isWindowsFirewallReachable() bool {
args := []string{"advfirewall", "show", "allprofiles", "state"}
netshCmd := GetSystem32Command("netsh")
netshCmd := wincmd.System32("netsh")
cmd := exec.Command(netshCmd, args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
@@ -117,23 +118,10 @@ func isWindowsFirewallReachable() bool {
func isFirewallRuleActive(ruleName string) bool {
args := []string{"advfirewall", "firewall", "show", "rule", "name=" + ruleName}
netshCmd := GetSystem32Command("netsh")
netshCmd := wincmd.System32("netsh")
cmd := exec.Command(netshCmd, args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
_, err := cmd.Output()
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"
}
-10
View File
@@ -486,16 +486,6 @@ func incrementalUpdate(oldChecksum uint16, oldBytes, newBytes []byte) uint16 {
return ^uint16(sum)
}
// AddDNATRule adds outbound DNAT rule for forwarding external traffic to NetBird network.
func (m *Manager) AddDNATRule(firewall.ForwardRule) (firewall.Rule, error) {
return nil, errNotSupported
}
// DeleteDNATRule deletes outbound DNAT rule.
func (m *Manager) DeleteDNATRule(firewall.Rule) error {
return errNotSupported
}
// addPortRedirection adds a port redirection rule.
func (m *Manager) addPortRedirection(targetIP netip.Addr, protocol gopacket.LayerType, originalPort, translatedPort uint16) error {
m.portDNATMutex.Lock()
+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 {
ipNets := make([]net.IPNet, len(prefixes))
for i, prefix := range prefixes {
normalized := normalizePrefix(prefix)
ipNets[i] = net.IPNet{
IP: prefix.Addr().AsSlice(), // Convert netip.Addr to net.IP
Mask: net.CIDRMask(prefix.Bits(), prefix.Addr().BitLen()), // Create subnet mask
IP: normalized.Addr().AsSlice(),
Mask: net.CIDRMask(normalized.Bits(), normalized.Addr().BitLen()),
}
}
return ipNets
+74 -34
View File
@@ -6,6 +6,7 @@ import (
"fmt"
"net"
"net/netip"
"slices"
"time"
log "github.com/sirupsen/logrus"
@@ -18,16 +19,22 @@ import (
type KernelConfigurer struct {
deviceName string
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 {
c := &KernelConfigurer{
deviceName: deviceName,
allowedIPs: newAllowedIPStore(),
}
c.statsCache = newStatsCache(statsCacheTTL, c.fetchStats)
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 {
log.Debugf("adding Wireguard private key")
key, err := wgtypes.ParseKey(privateKey)
@@ -46,6 +53,8 @@ func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error
if err != nil {
return fmt.Errorf(`received error "%w" while configuring interface %s with port %d`, err, c.deviceName, port)
}
c.allowedIPs.reset()
return nil
}
@@ -58,9 +67,20 @@ func (c *KernelConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, upda
}
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 {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
@@ -83,19 +103,23 @@ func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix,
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())
}
c.allowedIPs.add(peerKeyParsed, allowedIps)
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 {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return err
}
// Get the existing peer to preserve its allowed IPs
existingPeer, err := c.getPeer(c.deviceName, peerKey)
allowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
if err != nil {
return fmt.Errorf("get peer: %w", err)
return err
}
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 {
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{
PublicKey: peerKeyParsed,
AllowedIPs: existingPeer.AllowedIPs,
AllowedIPs: prefixesToIPNets(allowedIPs),
ReplaceAllowedIPs: true,
}
if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{reAddPeerCfg}}); err != nil {
c.allowedIPs.forget(peerKeyParsed)
return fmt.Errorf(
`error re-adding peer %s to interface %s with allowed IPs %v: %w`,
peerKey, c.deviceName, existingPeer.AllowedIPs, err,
"re-add peer %s to interface %s with allowed IPs %v: %w",
peerKey, c.deviceName, allowedIPs, err,
)
}
return nil
}
// RemovePeer removes a peer and forgets its allowed IPs after a successful device write.
func (c *KernelConfigurer) RemovePeer(peerKey string) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
@@ -142,15 +167,13 @@ func (c *KernelConfigurer) RemovePeer(peerKey string) error {
if err != nil {
return fmt.Errorf(`received error "%w" while removing peer %s from interface %s`, err, peerKey, c.deviceName)
}
c.allowedIPs.forget(peerKeyParsed)
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 {
ipNet := net.IPNet{
IP: allowedIP.Addr().AsSlice(),
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
}
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return err
@@ -159,7 +182,7 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix)
PublicKey: peerKeyParsed,
UpdateOnly: true,
ReplaceAllowedIPs: false,
AllowedIPs: []net.IPNet{ipNet},
AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}),
}
config := wgtypes.Config{
@@ -169,52 +192,69 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix)
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)
}
c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP})
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 {
ipNet := net.IPNet{
IP: allowedIP.Addr().AsSlice(),
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
}
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return fmt.Errorf("parse peer key: %w", err)
}
existingPeer, err := c.getPeer(c.deviceName, peerKey)
currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
if err != nil {
return fmt.Errorf("get peer: %w", err)
return err
}
newAllowedIPs := existingPeer.AllowedIPs
for i, existingAllowedIP := range existingPeer.AllowedIPs {
if existingAllowedIP.String() == ipNet.String() {
newAllowedIPs = append(existingPeer.AllowedIPs[:i], existingPeer.AllowedIPs[i+1:]...) //nolint:gocritic
break
}
idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP))
if idx < 0 {
return nil
}
newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1)
peer := wgtypes.PeerConfig{
PublicKey: peerKeyParsed,
UpdateOnly: true,
ReplaceAllowedIPs: true,
AllowedIPs: newAllowedIPs,
AllowedIPs: prefixesToIPNets(newAllowedIPs),
}
config := wgtypes.Config{
Peers: []wgtypes.PeerConfig{peer},
}
err = c.configure(config)
if err != nil {
if err := c.configure(config); err != nil {
return fmt.Errorf("remove allowed IP %s on interface %s: %w", allowedIP, c.deviceName, err)
}
c.allowedIPs.set(peerKeyParsed, newAllowedIPs)
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()
if err != nil {
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)
}
for _, peer := range wgDevice.Peers {
if peer.PublicKey.String() == peerPubKey {
if peer.PublicKey == peerPubKey {
return peer, nil
}
}
+120 -92
View File
@@ -8,6 +8,7 @@ import (
"net/netip"
"os"
"runtime"
"slices"
"strconv"
"strings"
"time"
@@ -41,31 +42,38 @@ type WGUSPConfigurer struct {
deviceName string
activityRecorder *bind.ActivityRecorder
statsCache *statsCache
allowedIPs *allowedIPStore
uapiListener net.Listener
}
// NewUSPConfigurer creates a userspace configurer and starts its UAPI listener.
func NewUSPConfigurer(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
wgCfg := &WGUSPConfigurer{
device: device,
deviceName: deviceName,
activityRecorder: activityRecorder,
allowedIPs: newAllowedIPStore(),
}
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
wgCfg.startUAPI()
return wgCfg
}
// NewUSPConfigurerNoUAPI creates a userspace configurer without a UAPI listener.
func NewUSPConfigurerNoUAPI(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
wgCfg := &WGUSPConfigurer{
device: device,
deviceName: deviceName,
activityRecorder: activityRecorder,
allowedIPs: newAllowedIPStore(),
}
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
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 {
log.Debugf("adding Wireguard private key")
key, err := wgtypes.ParseKey(privateKey)
@@ -80,7 +88,12 @@ func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error
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.
@@ -92,14 +105,38 @@ func (c *WGUSPConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, updat
}
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 {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
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{
PublicKey: peerKeyParsed,
ReplaceAllowedIPs: false,
@@ -119,47 +156,27 @@ func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix,
}
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.allowedIPs.add(peerKeyParsed, allowedIps)
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 {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return fmt.Errorf("parse peer key: %w", err)
}
ipcStr, err := c.device.IpcGet()
allowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
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{
PublicKey: peerKeyParsed,
Remove: true,
@@ -169,14 +186,13 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
Peers: []wgtypes.PeerConfig{peer},
}
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{
PublicKey: peerKeyParsed,
ReplaceAllowedIPs: true,
AllowedIPs: allowedIPs,
AllowedIPs: prefixesToIPNets(allowedIPs),
}
config = wgtypes.Config{
@@ -184,12 +200,15 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
}
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
}
// 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 {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
@@ -204,18 +223,17 @@ func (c *WGUSPConfigurer) RemovePeer(peerKey string) error {
config := wgtypes.Config{
Peers: []wgtypes.PeerConfig{peer},
}
ipcErr := c.device.IpcSet(toWgUserspaceString(config))
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()),
if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil {
return ipcErr
}
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)
if err != nil {
return err
@@ -224,79 +242,89 @@ func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) e
PublicKey: peerKeyParsed,
UpdateOnly: true,
ReplaceAllowedIPs: false,
AllowedIPs: []net.IPNet{ipNet},
AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}),
}
config := wgtypes.Config{
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 {
ipc, err := c.device.IpcGet()
if err != nil {
return err
}
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return fmt.Errorf("parse peer key: %w", err)
}
currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
if err != nil {
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{
PublicKey: peerKeyParsed,
UpdateOnly: 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{
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) {
@@ -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 {
GetProxy() wgproxy.Proxy
GetProxyPort() uint16
Free() error
}
@@ -81,12 +80,6 @@ func (w *WGIface) GetProxy() wgproxy.Proxy {
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.
func (w *WGIface) GetBind() device.EndpointManager {
w.mu.Lock()
-1
View File
@@ -52,7 +52,6 @@ func (f *fakeTunDevice) Close() error {
type fakeProxyFactory struct{}
func (fakeProxyFactory) GetProxy() wgproxy.Proxy { return nil }
func (fakeProxyFactory) GetProxyPort() uint16 { return 0 }
func (fakeProxyFactory) Free() error { return nil }
// TestWGIface_CloseReleasesMutexBeforeTunClose guards against a deadlock
+2 -15
View File
@@ -6,27 +6,14 @@ import (
"fmt"
"os/exec"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/wincmd"
)
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()
if err != nil {
return fmt.Errorf("failed to remove interface %s: %w - %s", w.Name(), err, out)
}
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()
}
// 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) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
addr := "100.64.0.1/8"
wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{
IFaceName: ifaceName,
@@ -127,10 +131,7 @@ func getIfaceAddrs(ifaceName string) ([]net.Addr, error) {
func Test_CreateInterface(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1)
wgIP := "10.99.99.1/32"
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{
IFaceName: ifaceName,
Address: wgaddr.MustParseWGAddress(wgIP),
@@ -170,10 +171,7 @@ func Test_Close(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
wgIP := "10.99.99.2/32"
wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{
IFaceName: ifaceName,
@@ -215,10 +213,7 @@ func TestRecreation(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
wgIP := "10.99.99.2/32"
wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{
IFaceName: ifaceName,
@@ -288,10 +283,7 @@ func Test_ConfigureInterface(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3)
wgIP := "10.99.99.5/30"
wgPort := 33100
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{
IFaceName: ifaceName,
Address: wgaddr.MustParseWGAddress(wgIP),
@@ -343,10 +335,7 @@ func Test_ConfigureInterface(t *testing.T) {
func Test_UpdatePeer(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
wgIP := "10.99.99.9/30"
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{
IFaceName: ifaceName,
@@ -413,10 +402,7 @@ func Test_UpdatePeer(t *testing.T) {
func Test_RemovePeer(t *testing.T) {
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
wgIP := "10.99.99.13/30"
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
opts := WGIFaceOpts{
IFaceName: ifaceName,
@@ -477,10 +463,7 @@ func Test_ConnectPeers(t *testing.T) {
peer2wgPort := 33200
keepAlive := 1 * time.Second
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
guid := fmt.Sprintf("{%s}", uuid.New().String())
device.CustomWindowsGUIDString = strings.ToLower(guid)
@@ -516,10 +499,7 @@ func Test_ConnectPeers(t *testing.T) {
guid = fmt.Sprintf("{%s}", uuid.New().String())
device.CustomWindowsGUIDString = strings.ToLower(guid)
newNet, err = stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet = stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
optsPeer2 := WGIFaceOpts{
IFaceName: peer2ifaceName,
@@ -568,11 +548,14 @@ func Test_ConnectPeers(t *testing.T) {
if err != nil {
t.Fatal(err)
}
// The peers use userspace WireGuard (stdnet transport). A tight busy-loop
// here starves the wireguard-go goroutines that process the handshake, so
// poll on a ticker instead and yield the CPU between checks. WireGuard also
// only retries a lost handshake initiation every REKEY_TIMEOUT (5s), which
// is why the overall wait can occasionally stretch to tens of seconds.
// On Linux with the kernel module both peers are kernel devices, elsewhere
// they run on wireguard-go. A tight busy-loop here would starve the
// wireguard-go goroutines that process the handshake, so poll on a ticker
// instead and yield the CPU between checks. WireGuard also only retries a
// 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
timeoutChannel := time.After(timeout)
ticker := time.NewTicker(500 * time.Millisecond)
@@ -590,13 +573,26 @@ func Test_ConnectPeers(t *testing.T) {
select {
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:
}
}
}
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) {
wg, err := wgctrl.New()
if err != nil {
+1 -4
View File
@@ -200,10 +200,7 @@ func (m *SingleSocketUDPMux) updateLocalAddresses() {
}
if len(networks) > 0 {
if m.params.Net == nil {
var err error
if m.params.Net, err = stdnet.NewNet(context.Background(), nil); err != nil {
m.params.Logger.Errorf("failed to get create network: %v", err)
}
m.params.Net = stdnet.NewNet(context.Background(), nil, nil)
}
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"
"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"
)
const (
envDisableKernelWGProxy = "NB_DISABLE_KERNEL_WG_PROXY"
// envDisableEBPFWGProxy is a deprecated alias for envDisableKernelWGProxy.
envDisableEBPFWGProxy = "NB_DISABLE_EBPF_WG_PROXY"
)
@@ -20,7 +22,7 @@ type KernelFactory struct {
wgPort int
mtu uint16
ebpfProxy *ebpf.WGEBPFProxy
loopbackProxy *loopback.Proxy
}
func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory {
@@ -29,55 +31,56 @@ func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory {
mtu: mtu,
}
if isEBPFDisabled() {
if isKernelProxyDisabled() {
log.Infof("WireGuard Proxy Factory will produce UDP proxy")
log.Infof("eBPF WireGuard proxy is disabled via %s environment variable", envDisableEBPFWGProxy)
return f
}
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, mtu)
if err := ebpfProxy.Listen(); err != nil {
loopbackProxy := loopback.NewProxy(wgPort, mtu)
if err := loopbackProxy.Listen(); err != nil {
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
}
log.Infof("WireGuard Proxy Factory will produce eBPF proxy")
f.ebpfProxy = ebpfProxy
log.Infof("WireGuard Proxy Factory will produce loopback proxy")
f.loopbackProxy = loopbackProxy
return f
}
func (w *KernelFactory) GetProxy() Proxy {
if w.ebpfProxy == nil {
if w.loopbackProxy == nil {
return udpProxy.NewWGUDPProxy(w.wgPort, w.mtu)
}
return ebpf.NewProxyWrapper(w.ebpfProxy)
}
// 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()
return loopback.NewProxyWrapper(w.loopbackProxy)
}
func (w *KernelFactory) Free() error {
if w.ebpfProxy == nil {
if w.loopbackProxy == nil {
return nil
}
return w.ebpfProxy.Free()
return w.loopbackProxy.Free()
}
func isEBPFDisabled() bool {
val := os.Getenv(envDisableEBPFWGProxy)
func isKernelProxyDisabled() bool {
env := envDisableKernelWGProxy
val := os.Getenv(env)
if val == "" {
env = envDisableEBPFWGProxy
val = os.Getenv(env)
}
if val == "" {
return false
}
disabled, err := strconv.ParseBool(val)
if err != nil {
log.Warnf("failed to parse %s: %v", envDisableEBPFWGProxy, err)
log.Warnf("failed to parse %s: %v", env, err)
return false
}
if disabled {
log.Infof("kernel WireGuard proxy is disabled via %s", env)
}
return disabled
}
-5
View File
@@ -24,11 +24,6 @@ func (w *USPFactory) GetProxy() Proxy {
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 {
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
package ebpf
package loopback
import (
"context"
@@ -8,6 +8,7 @@ import (
"fmt"
"io"
"net"
"net/netip"
"sync"
"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
type ProxyWrapper struct {
wgeBPFProxy *WGEBPFProxy
proxy *Proxy
remoteConn net.Conn
ctx context.Context
cancel context.CancelFunc
wgRelayedEndpointAddr *net.UDPAddr
peerAddr netip.Addr
headers *PacketHeaders
headerCurrentUsed *PacketHeaders
rawConn net.PacketConn
@@ -113,36 +115,44 @@ type ProxyWrapper struct {
closeListener *listener.CloseListener
}
func NewProxyWrapper(proxy *WGEBPFProxy) *ProxyWrapper {
func NewProxyWrapper(proxy *Proxy) *ProxyWrapper {
return &ProxyWrapper{
wgeBPFProxy: proxy,
proxy: proxy,
pausedCond: sync.NewCond(&sync.Mutex{}),
closeListener: listener.NewCloseListener(),
}
}
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 {
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 {
release()
return fmt.Errorf("create packet sender: %w", err)
}
// Check if required raw connection is available
if !headers.isIPv4 && p.wgeBPFProxy.rawConnIPv6 == nil {
if !headers.isIPv4 && p.proxy.rawConnIPv6 == nil {
release()
return errIPv6ConnNotAvailable
}
if headers.isIPv4 && p.wgeBPFProxy.rawConnIPv4 == nil {
if headers.isIPv4 && p.proxy.rawConnIPv4 == nil {
release()
return errIPv4ConnNotAvailable
}
p.remoteConn = remoteConn
p.ctx, p.cancel = context.WithCancel(ctx)
p.wgRelayedEndpointAddr = addr
p.peerAddr = peerAddr
p.headers = headers
p.rawConn = p.selectRawConn(headers)
return nil
@@ -193,18 +203,18 @@ func (p *ProxyWrapper) RedirectAs(endpoint *net.UDPAddr) {
return
}
header, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, endpoint)
header, err := NewPacketHeaders(p.proxy.localWGListenPort, endpoint)
if err != nil {
log.Errorf("failed to create packet headers: %s", err)
return
}
// Check if required raw connection is available
if !header.isIPv4 && p.wgeBPFProxy.rawConnIPv6 == nil {
if !header.isIPv4 && p.proxy.rawConnIPv6 == nil {
log.Error(errIPv6ConnNotAvailable)
return
}
if header.isIPv4 && p.wgeBPFProxy.rawConnIPv4 == nil {
if header.isIPv4 && p.proxy.rawConnIPv4 == nil {
log.Error(errIPv4ConnNotAvailable)
return
}
@@ -240,6 +250,10 @@ func (p *ProxyWrapper) CloseConn() error {
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.paused = false
p.pausedCond.Signal()
@@ -252,9 +266,9 @@ func (p *ProxyWrapper) CloseConn() error {
}
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 {
n, err := p.readFromRemote(ctx, buf)
if err != nil {
@@ -286,7 +300,7 @@ func (p *ProxyWrapper) readFromRemote(ctx context.Context, buf []byte) (int, err
}
p.closeListener.Notify()
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
}
@@ -314,7 +328,7 @@ func (p *ProxyWrapper) sendPkg(data []byte, header *PacketHeaders) error {
func (p *ProxyWrapper) selectRawConn(header *PacketHeaders) net.PacketConn {
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/wgaddr"
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"
)
func seedProxies() ([]proxyInstance, error) {
pl := make([]proxyInstance, 0)
ebpfProxy := ebpf.NewWGEBPFProxy(51831, 1280)
if err := ebpfProxy.Listen(); err != nil {
return nil, fmt.Errorf("failed to initialize ebpf proxy: %s", err)
loopbackProxy := loopback.NewProxy(51831, 1280)
if err := loopbackProxy.Listen(); err != nil {
return nil, fmt.Errorf("failed to initialize loopback proxy: %s", err)
}
pEbpf := proxyInstance{
name: "ebpf kernel proxy",
proxy: ebpf.NewProxyWrapper(ebpfProxy),
pLoopback := proxyInstance{
name: "loopback kernel proxy",
proxy: loopback.NewProxyWrapper(loopbackProxy),
wgPort: 51831,
closeFn: ebpfProxy.Free,
closeFn: loopbackProxy.Free,
}
pl = append(pl, pEbpf)
pl = append(pl, pLoopback)
pUDP := proxyInstance{
name: "udp kernel proxy",
@@ -42,18 +42,18 @@ func seedProxies() ([]proxyInstance, error) {
func seedProxyForProxyCloseByRemoteConn() ([]proxyInstance, error) {
pl := make([]proxyInstance, 0)
ebpfProxy := ebpf.NewWGEBPFProxy(51831, 1280)
if err := ebpfProxy.Listen(); err != nil {
return nil, fmt.Errorf("failed to initialize ebpf proxy: %s", err)
loopbackProxy := loopback.NewProxy(51831, 1280)
if err := loopbackProxy.Listen(); err != nil {
return nil, fmt.Errorf("failed to initialize loopback proxy: %s", err)
}
pEbpf := proxyInstance{
name: "ebpf kernel proxy",
proxy: ebpf.NewProxyWrapper(ebpfProxy),
pLoopback := proxyInstance{
name: "loopback kernel proxy",
proxy: loopback.NewProxyWrapper(loopbackProxy),
wgPort: 51831,
closeFn: ebpfProxy.Free,
closeFn: loopbackProxy.Free,
}
pl = append(pl, pEbpf)
pl = append(pl, pLoopback)
pUDP := proxyInstance{
name: "udp kernel proxy",
+23 -23
View File
@@ -8,7 +8,7 @@ import (
"testing"
"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"
)
@@ -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
func TestRedirectAs_eBPF_IPv4(t *testing.T) {
// TestRedirectAs_Loopback_IPv4 tests RedirectAs with the loopback proxy using IPv4 addresses
func TestRedirectAs_Loopback_IPv4(t *testing.T) {
wgPort := 51850
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280)
if err := ebpfProxy.Listen(); err != nil {
t.Fatalf("failed to initialize ebpf proxy: %v", err)
loopbackProxy := loopback.NewProxy(wgPort, 1280)
if err := loopbackProxy.Listen(); err != nil {
t.Fatalf("failed to initialize loopback proxy: %v", err)
}
defer func() {
if err := ebpfProxy.Free(); err != nil {
t.Errorf("failed to free ebpf proxy: %v", err)
if err := loopbackProxy.Free(); err != nil {
t.Errorf("failed to free loopback proxy: %v", err)
}
}()
proxy := ebpf.NewProxyWrapper(ebpfProxy)
proxy := loopback.NewProxyWrapper(loopbackProxy)
// NetBird UDP address of the remote peer
nbAddr := &net.UDPAddr{
@@ -227,20 +227,20 @@ func TestRedirectAs_eBPF_IPv4(t *testing.T) {
testRedirectAs(t, proxy, wgPort, nbAddr, p2pEndpoint)
}
// TestRedirectAs_eBPF_IPv6 tests RedirectAs with eBPF proxy using IPv6 addresses
func TestRedirectAs_eBPF_IPv6(t *testing.T) {
// TestRedirectAs_Loopback_IPv6 tests RedirectAs with the loopback proxy using IPv6 addresses
func TestRedirectAs_Loopback_IPv6(t *testing.T) {
wgPort := 51851
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280)
if err := ebpfProxy.Listen(); err != nil {
t.Fatalf("failed to initialize ebpf proxy: %v", err)
loopbackProxy := loopback.NewProxy(wgPort, 1280)
if err := loopbackProxy.Listen(); err != nil {
t.Fatalf("failed to initialize loopback proxy: %v", err)
}
defer func() {
if err := ebpfProxy.Free(); err != nil {
t.Errorf("failed to free ebpf proxy: %v", err)
if err := loopbackProxy.Free(); err != nil {
t.Errorf("failed to free loopback proxy: %v", err)
}
}()
proxy := ebpf.NewProxyWrapper(ebpfProxy)
proxy := loopback.NewProxyWrapper(loopbackProxy)
// NetBird UDP address of the remote peer
nbAddr := &net.UDPAddr{
@@ -259,17 +259,17 @@ func TestRedirectAs_eBPF_IPv6(t *testing.T) {
// TestRedirectAs_Multiple_Switches tests switching between multiple endpoints
func TestRedirectAs_Multiple_Switches(t *testing.T) {
wgPort := 51856
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280)
if err := ebpfProxy.Listen(); err != nil {
t.Fatalf("failed to initialize ebpf proxy: %v", err)
loopbackProxy := loopback.NewProxy(wgPort, 1280)
if err := loopbackProxy.Listen(); err != nil {
t.Fatalf("failed to initialize loopback proxy: %v", err)
}
defer func() {
if err := ebpfProxy.Free(); err != nil {
t.Errorf("failed to free ebpf proxy: %v", err)
if err := loopbackProxy.Free(); err != nil {
t.Errorf("failed to free loopback proxy: %v", err)
}
}()
proxy := ebpf.NewProxyWrapper(ebpfProxy)
proxy := loopback.NewProxyWrapper(loopbackProxy)
ctx := context.Background()
+113
View File
@@ -0,0 +1,113 @@
package auth
import (
"encoding/base64"
"encoding/json"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestTokenInfoMatchesAccount(t *testing.T) {
tests := []struct {
name string
token TokenInfo
hint string
match bool
}{
{
name: "same account",
token: TokenInfo{EmailClaim: "user@example.com"},
hint: "user@example.com",
match: true,
},
{
name: "different account",
token: TokenInfo{EmailClaim: "other@example.com"},
hint: "user@example.com",
match: false,
},
{
name: "case differences are the same account",
token: TokenInfo{EmailClaim: "User@Example.com"},
hint: "user@example.com",
match: true,
},
{
name: "no hint leaves the choice to the IdP",
token: TokenInfo{EmailClaim: "other@example.com"},
hint: "",
match: true,
},
{
name: "token without an email claim is not judged",
token: TokenInfo{EmailClaim: ""},
hint: "user@example.com",
match: true,
},
{
name: "name fallback does not trigger matching",
token: TokenInfo{Email: "Some One"},
hint: "user@example.com",
match: true,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
assert.Equal(t, tc.match, tc.token.MatchesAccount(tc.hint))
})
}
}
func TestParseEmailFromIDToken(t *testing.T) {
tests := []struct {
name string
claims map[string]interface{}
wantValue string
wantFromEmail bool
wantErr bool
}{
{
name: "email claim",
claims: map[string]interface{}{"email": "user@example.com", "name": "Some One"},
wantValue: "user@example.com",
wantFromEmail: true,
},
{
name: "name fallback",
claims: map[string]interface{}{"name": "Some One"},
wantValue: "Some One",
},
{
name: "neither claim present",
claims: map[string]interface{}{"sub": "abc"},
wantErr: true,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
value, fromEmailClaim, err := parseEmailFromIDToken(idTokenWithClaims(t, tc.claims))
if tc.wantErr {
require.Error(t, err)
return
}
require.NoError(t, err)
assert.Equal(t, tc.wantValue, value)
assert.Equal(t, tc.wantFromEmail, fromEmailClaim)
})
}
}
func TestRetryFlowForAccountUnsupportedFlow(t *testing.T) {
assert.Nil(t, RetryFlowForAccount(&DeviceAuthorizationFlow{}))
}
func idTokenWithClaims(t *testing.T, claims map[string]interface{}) string {
t.Helper()
payload, err := json.Marshal(claims)
require.NoError(t, err)
return "header." + base64.RawURLEncoding.EncodeToString(payload) + ".signature"
}
+11 -7
View File
@@ -103,7 +103,7 @@ func (a *Auth) IsSSOSupported(ctx context.Context) (bool, error) {
err := a.withRetry(ctx, func(client *mgm.GrpcClient) error {
// Try PKCE flow first
_, err := a.getPKCEFlow(client)
_, err := a.getPKCEFlow(client, false)
if err == nil {
supportsSSO = true
return nil
@@ -136,9 +136,13 @@ func (a *Auth) IsSSOSupported(ctx context.Context) (bool, error) {
return supportsSSO, err
}
// GetOAuthFlow returns an OAuth flow (PKCE or Device) using the existing management connection
// This avoids creating a new connection to the management server
func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool, hint string) (OAuthFlow, error) {
// GetOAuthFlow returns an OAuth flow (PKCE or Device) using the existing management connection.
// This avoids creating a new connection to the management server.
//
// sessionExtend marks the flow as renewing an existing peer's session rather than
// logging one in; the server needs it to rule out a silent authorization that the
// IdP could answer from another account. See PKCEAuthorizationFlowRequest.
func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool, sessionExtend bool, hint string) (OAuthFlow, error) {
var flow OAuthFlow
err := a.withRetry(ctx, func(client *mgm.GrpcClient) error {
@@ -153,7 +157,7 @@ func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool, hint stri
}
// Try PKCE flow first
pkceFlow, err := a.getPKCEFlow(client)
pkceFlow, err := a.getPKCEFlow(client, sessionExtend)
if err != nil {
// If PKCE not supported, try Device flow
if s, ok := status.FromError(err); ok && (s.Code() == codes.NotFound || s.Code() == codes.Unimplemented) {
@@ -240,8 +244,8 @@ func (a *Auth) Login(ctx context.Context, setupKey string, jwtToken string) (err
}
// getPKCEFlow retrieves PKCE authorization flow configuration and creates a flow instance
func (a *Auth) getPKCEFlow(client *mgm.GrpcClient) (*PKCEAuthorizationFlow, error) {
protoFlow, err := client.GetPKCEAuthorizationFlow()
func (a *Auth) getPKCEFlow(client *mgm.GrpcClient, sessionExtend bool) (*PKCEAuthorizationFlow, error) {
protoFlow, err := client.GetPKCEAuthorizationFlow(sessionExtend)
if err != nil {
if s, ok := status.FromError(err); ok && s.Code() == codes.NotFound {
log.Warnf("server couldn't find pkce flow, contact admin: %v", err)
+4 -1
View File
@@ -308,10 +308,13 @@ func (d *DeviceAuthorizationFlow) WaitToken(ctx context.Context, info AuthFlowIn
// callers store to send back as the login_hint. Without it a client
// driven through the device flow — Android TV and tvOS — never binds
// an account to its profile and every later login goes out blind.
if email, err := parseEmailFromIDToken(tokenInfo.IDToken); err != nil {
if email, fromEmailClaim, err := parseEmailFromIDToken(tokenInfo.IDToken); err != nil {
log.Warnf("failed to parse email from ID token: %v", err)
} else {
tokenInfo.Email = email
if fromEmailClaim {
tokenInfo.EmailClaim = email
}
}
log.Infof("device flow: user authorization confirmed after %d polls in %s", polls, time.Since(start).Round(time.Second))

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