diff --git a/.coderabbit.yaml b/.coderabbit.yaml index 85ed5cd3b..a1336cdab 100644 --- a/.coderabbit.yaml +++ b/.coderabbit.yaml @@ -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 diff --git a/.githooks/commit-msg b/.githooks/commit-msg new file mode 100755 index 000000000..e5699bd9c --- /dev/null +++ b/.githooks/commit-msg @@ -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 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 diff --git a/.github/dependabot.yml b/.github/dependabot.yml index 647e04936..ded77ec58 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -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" diff --git a/.github/scripts/test-homebrew-cask.sh b/.github/scripts/test-homebrew-cask.sh new file mode 100755 index 000000000..9c1041bdd --- /dev/null +++ b/.github/scripts/test-homebrew-cask.sh @@ -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" < '$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::" diff --git a/.github/workflows/check-license-dependencies.yml b/.github/workflows/check-license-dependencies.yml index 81d293e4f..27c59f8d6 100644 --- a/.github/workflows/check-license-dependencies.yml +++ b/.github/workflows/check-license-dependencies.yml @@ -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" diff --git a/.github/workflows/docs-ack.yml b/.github/workflows/docs-ack.yml index 7e34e2f8a..edd1eebff 100644 --- a/.github/workflows/docs-ack.yml +++ b/.github/workflows/docs-ack.yml @@ -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 diff --git a/.github/workflows/frontend-ui.yml b/.github/workflows/frontend-ui.yml index 014c5c2ae..2ad43c581 100644 --- a/.github/workflows/frontend-ui.yml +++ b/.github/workflows/frontend-ui.yml @@ -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') }} diff --git a/.github/workflows/golang-test-darwin.yml b/.github/workflows/golang-test-darwin.yml index c17d8e775..dec57dc80 100644 --- a/.github/workflows/golang-test-darwin.yml +++ b/.github/workflows/golang-test-darwin.yml @@ -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 diff --git a/.github/workflows/golang-test-freebsd.yml b/.github/workflows/golang-test-freebsd.yml index 65c39147a..7bd48e3d0 100644 --- a/.github/workflows/golang-test-freebsd.yml +++ b/.github/workflows/golang-test-freebsd.yml @@ -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 diff --git a/.github/workflows/golang-test-linux.yml b/.github/workflows/golang-test-linux.yml index f24dfbe9d..dd6b9fe68 100644 --- a/.github/workflows/golang-test-linux.yml +++ b/.github/workflows/golang-test-linux.yml @@ -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' diff --git a/.github/workflows/golang-test-windows.yml b/.github/workflows/golang-test-windows.yml index fb7b745d2..49be3af58 100644 --- a/.github/workflows/golang-test-windows.yml +++ b/.github/workflows/golang-test-windows.yml @@ -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 diff --git a/.github/workflows/golangci-lint.yml b/.github/workflows/golangci-lint.yml index 586e1235b..843177b88 100644 --- a/.github/workflows/golangci-lint.yml +++ b/.github/workflows/golangci-lint.yml @@ -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 }} diff --git a/.github/workflows/mobile-build-validation.yml b/.github/workflows/mobile-build-validation.yml new file mode 100644 index 000000000..613a39b3d --- /dev/null +++ b/.github/workflows/mobile-build-validation.yml @@ -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/... diff --git a/.github/workflows/pr-title-check.yml b/.github/workflows/pr-title-check.yml index 24d81b50f..5b769e2f7 100644 --- a/.github/workflows/pr-title-check.yml +++ b/.github/workflows/pr-title-check.yml @@ -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 diff --git a/.github/workflows/redhat-certify.yml b/.github/workflows/redhat-certify.yml new file mode 100644 index 000000000..7c26c79fb --- /dev/null +++ b/.github/workflows/redhat-certify.yml @@ -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_ +# 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 diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index c1bbe9c44..a79357505 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -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__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 diff --git a/.github/workflows/test-homebrew-cask.yml b/.github/workflows/test-homebrew-cask.yml new file mode 100644 index 000000000..a75951425 --- /dev/null +++ b/.github/workflows/test-homebrew-cask.yml @@ -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 diff --git a/.github/workflows/ui-translations.yml b/.github/workflows/ui-translations.yml index 7d3b12f2d..9ac524495 100644 --- a/.github/workflows/ui-translations.yml +++ b/.github/workflows/ui-translations.yml @@ -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 diff --git a/.gitignore b/.gitignore index 305f3cb50..5c01f6e60 100644 --- a/.gitignore +++ b/.gitignore @@ -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 diff --git a/.goreleaser.yaml b/.goreleaser.yaml index 759acb725..6ab9da749 100644 --- a/.goreleaser.yaml +++ b/.goreleaser.yaml @@ -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 @@ -223,23 +275,72 @@ nfpms: postinstall: "release_files/post_install.sh" preremove: "release_files/pre_remove.sh" - - maintainer: Netbird + - &netbird_rpm + maintainer: Netbird 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 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 diff --git a/AGENTS.md b/AGENTS.md index 5497acb15..3838a7913 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -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/... diff --git a/CLAUDE.md b/CLAUDE.md index 764f406be..72681748e 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -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 diff --git a/LICENSE b/LICENSE index d922f155a..cea6f8f0b 100644 --- a/LICENSE +++ b/LICENSE @@ -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 diff --git a/Makefile b/Makefile index 0a4fad2f2..26c5b932e 100644 --- a/Makefile +++ b/Makefile @@ -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. diff --git a/README.md b/README.md index 336332043..4dbfc7bd0 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/SUPPORT.md b/SUPPORT.md new file mode 100644 index 000000000..fade37286 --- /dev/null +++ b/SUPPORT.md @@ -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) diff --git a/agent-network/README.md b/agent-network/README.md index 029ada299..35b9c6668 100644 --- a/agent-network/README.md +++ b/agent-network/README.md @@ -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, diff --git a/base62/base62.go b/base62/base62.go index efafbc768..1a02e98e2 100644 --- a/base62/base62.go +++ b/base62/base62.go @@ -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) -} diff --git a/base62/base62_test.go b/base62/base62_test.go index 00da2124a..f2ad06d6f 100644 --- a/base62/base62_test.go +++ b/base62/base62_test.go @@ -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) + } +} diff --git a/client/Dockerfile-rootless.ubi b/client/Dockerfile-rootless.ubi new file mode 100644 index 000000000..4701728c1 --- /dev/null +++ b/client/Dockerfile-rootless.ubi @@ -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 " \ + 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"] diff --git a/client/android/client.go b/client/android/client.go index 5bd0d1e10..9705db8e0 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -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 } } diff --git a/client/android/client_mdm.go b/client/android/client_mdm.go new file mode 100644 index 000000000..d043b85d3 --- /dev/null +++ b/client/android/client_mdm.go @@ -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 +} diff --git a/client/android/login.go b/client/android/login.go index 3742e01a5..b9ec21b39 100644 --- a/client/android/login.go +++ b/client/android/login.go @@ -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. diff --git a/client/android/login_test.go b/client/android/login_test.go index b04790f6b..130a846fc 100644 --- a/client/android/login_test.go +++ b/client/android/login_test.go @@ -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) } diff --git a/client/android/mdm.go b/client/android/mdm.go new file mode 100644 index 000000000..617d8f7cb --- /dev/null +++ b/client/android/mdm.go @@ -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) +} diff --git a/client/android/preferences.go b/client/android/preferences.go index d90365518..3623de23f 100644 --- a/client/android/preferences.go +++ b/client/android/preferences.go @@ -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 } diff --git a/client/android/preferences_test.go b/client/android/preferences_test.go index 2bbccef86..d9f5b1918 100644 --- a/client/android/preferences_test.go +++ b/client/android/preferences_test.go @@ -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") } } diff --git a/client/android/profile_manager.go b/client/android/profile_manager.go index 557c837a7..4bc60c453 100644 --- a/client/android/profile_manager.go +++ b/client/android/profile_manager.go @@ -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) { diff --git a/client/android/session.go b/client/android/session.go index d5da09c93..b2de8dadb 100644 --- a/client/android/session.go +++ b/client/android/session.go @@ -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) } diff --git a/client/android/ssh_client.go b/client/android/ssh_client.go index 2822b6539..e05249e98 100644 --- a/client/android/ssh_client.go +++ b/client/android/ssh_client.go @@ -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) } diff --git a/client/cmd/debug.go b/client/cmd/debug.go index 98fe53626..c4d5ad6d7 100644 --- a/client/cmd/debug.go +++ b/client/cmd/debug.go @@ -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") } diff --git a/client/cmd/debug_cpu.go b/client/cmd/debug_cpu.go new file mode 100644 index 000000000..a01b845cf --- /dev/null +++ b/client/cmd/debug_cpu.go @@ -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) +} diff --git a/client/cmd/debug_cpu_test.go b/client/cmd/debug_cpu_test.go new file mode 100644 index 000000000..85fffd462 --- /dev/null +++ b/client/cmd/debug_cpu_test.go @@ -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 ` 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") +} diff --git a/client/cmd/forwarding_rules.go b/client/cmd/forwarding_rules.go deleted file mode 100644 index b3052746a..000000000 --- a/client/cmd/forwarding_rules.go +++ /dev/null @@ -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" - } -} diff --git a/client/cmd/login.go b/client/cmd/login.go index 4e08334eb..764fdd5b1 100644 --- a/client/cmd/login.go +++ b/client/cmd/login.go @@ -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) diff --git a/client/cmd/root.go b/client/cmd/root.go index be6479440..4525a9bd6 100644 --- a/client/cmd/root.go +++ b/client/cmd/root.go @@ -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) { diff --git a/client/cmd/service.go b/client/cmd/service.go index 7410d60ea..2a558e6d5 100644 --- a/client/cmd/service.go +++ b/client/cmd/service.go @@ -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) + }) +} diff --git a/client/cmd/service_params.go b/client/cmd/service_params.go index 750b22ae6..6e2dbec40 100644 --- a/client/cmd/service_params.go +++ b/client/cmd/service_params.go @@ -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)) diff --git a/client/cmd/service_params_test.go b/client/cmd/service_params_test.go index 94f98a0ce..1f83374cb 100644 --- a/client/cmd/service_params_test.go +++ b/client/cmd/service_params_test.go @@ -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 }) diff --git a/client/cmd/service_params_trust_test.go b/client/cmd/service_params_trust_test.go new file mode 100644 index 000000000..1cf564445 --- /dev/null +++ b/client/cmd/service_params_trust_test.go @@ -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) +} diff --git a/client/cmd/testutil_test.go b/client/cmd/testutil_test.go index 328a15454..46bf31837 100644 --- a/client/cmd/testutil_test.go +++ b/client/cmd/testutil_test.go @@ -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) } diff --git a/client/cmd/up.go b/client/cmd/up.go index 2e53224df..120a25595 100644 --- a/client/cmd/up.go +++ b/client/cmd/up.go @@ -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 diff --git a/client/cmd/up_setconfig_refusal_test.go b/client/cmd/up_setconfig_refusal_test.go new file mode 100644 index 000000000..fdf580102 --- /dev/null +++ b/client/cmd/up_setconfig_refusal_test.go @@ -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)) + }) + } +} diff --git a/client/embed/embed.go b/client/embed/embed.go index 5a3d11f24..5a3d540ec 100644 --- a/client/embed/embed.go +++ b/client/embed/embed.go @@ -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 diff --git a/client/embed/embed_test.go b/client/embed/embed_test.go index 4ff5c9978..a818af055 100644 --- a/client/embed/embed_test.go +++ b/client/embed/embed_test.go @@ -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) diff --git a/client/firewall/iptables/dnat_linux.go b/client/firewall/iptables/dnat_linux.go index eca8386c0..f118c9dfe 100644 --- a/client/firewall/iptables/dnat_linux.go +++ b/client/firewall/iptables/dnat_linux.go @@ -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)) diff --git a/client/firewall/iptables/dnat_refcount_linux_test.go b/client/firewall/iptables/dnat_refcount_linux_test.go deleted file mode 100644 index 40ebc6cc3..000000000 --- a/client/firewall/iptables/dnat_refcount_linux_test.go +++ /dev/null @@ -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") -} diff --git a/client/firewall/iptables/family_linux.go b/client/firewall/iptables/family_linux.go index c5ed8cc20..2ac860a0a 100644 --- a/client/firewall/iptables/family_linux.go +++ b/client/firewall/iptables/family_linux.go @@ -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. diff --git a/client/firewall/iptables/filter_linux.go b/client/firewall/iptables/filter_linux.go index dc606da2d..30cd81018 100644 --- a/client/firewall/iptables/filter_linux.go +++ b/client/firewall/iptables/filter_linux.go @@ -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 -// "_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 diff --git a/client/firewall/iptables/manager_linux.go b/client/firewall/iptables/manager_linux.go index 49b88f1ea..a566909c8 100644 --- a/client/firewall/iptables/manager_linux.go +++ b/client/firewall/iptables/manager_linux.go @@ -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"} } diff --git a/client/firewall/iptables/manager_linux_test.go b/client/firewall/iptables/manager_linux_test.go index 9f53352e1..8435bf6a5 100644 --- a/client/firewall/iptables/manager_linux_test.go +++ b/client/firewall/iptables/manager_linux_test.go @@ -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. diff --git a/client/firewall/manager/firewall.go b/client/firewall/manager/firewall.go index 97a94d0f5..f8de1e2b5 100644 --- a/client/firewall/manager/firewall.go +++ b/client/firewall/manager/firewall.go @@ -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. diff --git a/client/firewall/manager/forward_rule.go b/client/firewall/manager/forward_rule.go deleted file mode 100644 index c2e9e5c60..000000000 --- a/client/firewall/manager/forward_rule.go +++ /dev/null @@ -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()) -} diff --git a/client/firewall/nftables/dnat_linux.go b/client/firewall/nftables/dnat_linux.go index 8eae694a2..c179d60cc 100644 --- a/client/firewall/nftables/dnat_linux.go +++ b/client/firewall/nftables/dnat_linux.go @@ -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 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)) diff --git a/client/firewall/nftables/dnat_refcount_linux_test.go b/client/firewall/nftables/dnat_refcount_linux_test.go deleted file mode 100644 index cdc24e77f..000000000 --- a/client/firewall/nftables/dnat_refcount_linux_test.go +++ /dev/null @@ -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") -} diff --git a/client/firewall/nftables/family_linux.go b/client/firewall/nftables/family_linux.go index 7a5df3ed7..4169c9d2d 100644 --- a/client/firewall/nftables/family_linux.go +++ b/client/firewall/nftables/family_linux.go @@ -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) } diff --git a/client/firewall/nftables/filter_linux.go b/client/firewall/nftables/filter_linux.go index ebd238063..bb3ac1dfe 100644 --- a/client/firewall/nftables/filter_linux.go +++ b/client/firewall/nftables/filter_linux.go @@ -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. diff --git a/client/firewall/nftables/manager_linux.go b/client/firewall/nftables/manager_linux.go index dbd5e4fa2..75405e213 100644 --- a/client/firewall/nftables/manager_linux.go +++ b/client/firewall/nftables/manager_linux.go @@ -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) } diff --git a/client/firewall/nftables/manager_linux_test.go b/client/firewall/nftables/manager_linux_test.go index 0ca56409e..4d6eec3c1 100644 --- a/client/firewall/nftables/manager_linux_test.go +++ b/client/firewall/nftables/manager_linux_test.go @@ -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) diff --git a/client/firewall/nftables/routing_linux.go b/client/firewall/nftables/routing_linux.go index 4115c94bd..e98471e8f 100644 --- a/client/firewall/nftables/routing_linux.go +++ b/client/firewall/nftables/routing_linux.go @@ -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) diff --git a/client/firewall/uspfilter/filter.go b/client/firewall/uspfilter/filter.go index 5e1366c1f..0c73400f1 100644 --- a/client/firewall/uspfilter/filter.go +++ b/client/firewall/uspfilter/filter.go @@ -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 { diff --git a/client/firewall/uspfilter/interface_allower_windows.go b/client/firewall/uspfilter/interface_allower_windows.go index 7f525e28c..4cd0fe969 100644 --- a/client/firewall/uspfilter/interface_allower_windows.go +++ b/client/firewall/uspfilter/interface_allower_windows.go @@ -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" -} diff --git a/client/firewall/uspfilter/nat.go b/client/firewall/uspfilter/nat.go index 06312aabf..49c26766a 100644 --- a/client/firewall/uspfilter/nat.go +++ b/client/firewall/uspfilter/nat.go @@ -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() diff --git a/client/iface/configurer/allowedips.go b/client/iface/configurer/allowedips.go new file mode 100644 index 000000000..193197d4a --- /dev/null +++ b/client/iface/configurer/allowedips.go @@ -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 +} diff --git a/client/iface/configurer/allowedips_test.go b/client/iface/configurer/allowedips_test.go new file mode 100644 index 000000000..1d272d8c4 --- /dev/null +++ b/client/iface/configurer/allowedips_test.go @@ -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) + } +} diff --git a/client/iface/configurer/common.go b/client/iface/configurer/common.go index 10162d703..40f8209e9 100644 --- a/client/iface/configurer/common.go +++ b/client/iface/configurer/common.go @@ -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 diff --git a/client/iface/configurer/kernel_unix.go b/client/iface/configurer/kernel_unix.go index da69c2a35..3a95249c1 100644 --- a/client/iface/configurer/kernel_unix.go +++ b/client/iface/configurer/kernel_unix.go @@ -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 } } diff --git a/client/iface/configurer/usp.go b/client/iface/configurer/usp.go index 2be1b861e..334d99369 100644 --- a/client/iface/configurer/usp.go +++ b/client/iface/configurer/usp.go @@ -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) { diff --git a/client/iface/configurer/usp_allowedips_test.go b/client/iface/configurer/usp_allowedips_test.go new file mode 100644 index 000000000..fba0ca546 --- /dev/null +++ b/client/iface/configurer/usp_allowedips_test.go @@ -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 +} diff --git a/client/iface/iface.go b/client/iface/iface.go index 247f421a2..f6006fa87 100644 --- a/client/iface/iface.go +++ b/client/iface/iface.go @@ -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() diff --git a/client/iface/iface_close_test.go b/client/iface/iface_close_test.go index 171e15d0a..ea3115ec0 100644 --- a/client/iface/iface_close_test.go +++ b/client/iface/iface_close_test.go @@ -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 diff --git a/client/iface/iface_destroy_windows.go b/client/iface/iface_destroy_windows.go index 0bfa4e211..54c0014c4 100644 --- a/client/iface/iface_destroy_windows.go +++ b/client/iface/iface_destroy_windows.go @@ -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" -} diff --git a/client/iface/iface_test.go b/client/iface/iface_test.go index 89c8cd16e..690a0f66f 100644 --- a/client/iface/iface_test.go +++ b/client/iface/iface_test.go @@ -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 { diff --git a/client/iface/udpmux/mux.go b/client/iface/udpmux/mux.go index c5d2de4a5..3aa0b4d88 100644 --- a/client/iface/udpmux/mux.go +++ b/client/iface/udpmux/mux.go @@ -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) diff --git a/client/iface/wgproxy/ebpf/portlookup.go b/client/iface/wgproxy/ebpf/portlookup.go deleted file mode 100644 index fce8f1507..000000000 --- a/client/iface/wgproxy/ebpf/portlookup.go +++ /dev/null @@ -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 -} diff --git a/client/iface/wgproxy/ebpf/portlookup_test.go b/client/iface/wgproxy/ebpf/portlookup_test.go deleted file mode 100644 index a2e92fc79..000000000 --- a/client/iface/wgproxy/ebpf/portlookup_test.go +++ /dev/null @@ -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 -} diff --git a/client/iface/wgproxy/ebpf/proxy.go b/client/iface/wgproxy/ebpf/proxy.go deleted file mode 100644 index 91c741c0d..000000000 --- a/client/iface/wgproxy/ebpf/proxy.go +++ /dev/null @@ -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 -} diff --git a/client/iface/wgproxy/ebpf/proxy_test.go b/client/iface/wgproxy/ebpf/proxy_test.go deleted file mode 100644 index 228c06c9b..000000000 --- a/client/iface/wgproxy/ebpf/proxy_test.go +++ /dev/null @@ -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") - } -} diff --git a/client/iface/wgproxy/factory_kernel.go b/client/iface/wgproxy/factory_kernel.go index 7821df3de..0b2329b96 100644 --- a/client/iface/wgproxy/factory_kernel.go +++ b/client/iface/wgproxy/factory_kernel.go @@ -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 } diff --git a/client/iface/wgproxy/factory_usp.go b/client/iface/wgproxy/factory_usp.go index bbd67e076..a1b1c34d7 100644 --- a/client/iface/wgproxy/factory_usp.go +++ b/client/iface/wgproxy/factory_usp.go @@ -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 } diff --git a/client/iface/wgproxy/loopback/addr.go b/client/iface/wgproxy/loopback/addr.go new file mode 100644 index 000000000..52feee295 --- /dev/null +++ b/client/iface/wgproxy/loopback/addr.go @@ -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 +} diff --git a/client/iface/wgproxy/loopback/addr_test.go b/client/iface/wgproxy/loopback/addr_test.go new file mode 100644 index 000000000..3755b7269 --- /dev/null +++ b/client/iface/wgproxy/loopback/addr_test.go @@ -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") + } +} diff --git a/client/iface/wgproxy/loopback/proxy.go b/client/iface/wgproxy/loopback/proxy.go new file mode 100644 index 000000000..f9364c766 --- /dev/null +++ b/client/iface/wgproxy/loopback/proxy.go @@ -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) +} diff --git a/client/iface/wgproxy/loopback/proxy_privileged_test.go b/client/iface/wgproxy/loopback/proxy_privileged_test.go new file mode 100644 index 000000000..6314fe2de --- /dev/null +++ b/client/iface/wgproxy/loopback/proxy_privileged_test.go @@ -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) + } +} diff --git a/client/iface/wgproxy/ebpf/wrapper.go b/client/iface/wgproxy/loopback/wrapper.go similarity index 84% rename from client/iface/wgproxy/ebpf/wrapper.go rename to client/iface/wgproxy/loopback/wrapper.go index f75e21aa6..a9cb1ab59 100644 --- a/client/iface/wgproxy/ebpf/wrapper.go +++ b/client/iface/wgproxy/loopback/wrapper.go @@ -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 } diff --git a/client/iface/wgproxy/proxy_linux_test.go b/client/iface/wgproxy/proxy_linux_test.go index e34dd3b6b..88d4588a5 100644 --- a/client/iface/wgproxy/proxy_linux_test.go +++ b/client/iface/wgproxy/proxy_linux_test.go @@ -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", diff --git a/client/iface/wgproxy/redirect_test.go b/client/iface/wgproxy/redirect_test.go index f0d59cc64..47f571f2b 100644 --- a/client/iface/wgproxy/redirect_test.go +++ b/client/iface/wgproxy/redirect_test.go @@ -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() diff --git a/client/internal/auth/account_match_test.go b/client/internal/auth/account_match_test.go new file mode 100644 index 000000000..b879cf9ab --- /dev/null +++ b/client/internal/auth/account_match_test.go @@ -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" +} diff --git a/client/internal/auth/auth.go b/client/internal/auth/auth.go index 939df3a21..7c156b5fd 100644 --- a/client/internal/auth/auth.go +++ b/client/internal/auth/auth.go @@ -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) diff --git a/client/internal/auth/device_flow.go b/client/internal/auth/device_flow.go index 3592e589d..c2b8d34b8 100644 --- a/client/internal/auth/device_flow.go +++ b/client/internal/auth/device_flow.go @@ -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)) diff --git a/client/internal/auth/oauth.go b/client/internal/auth/oauth.go index 91329c98b..54ecb5d86 100644 --- a/client/internal/auth/oauth.go +++ b/client/internal/auth/oauth.go @@ -5,6 +5,7 @@ import ( "fmt" "net/http" "runtime" + "strings" log "github.com/sirupsen/logrus" "google.golang.org/grpc/codes" @@ -25,6 +26,14 @@ type HTTPClient interface { Do(req *http.Request) (*http.Response, error) } +// accountPromptForcer is implemented by the PKCE flow only. The device code +// flow has no equivalent: RFC 8628 defines no prompt parameter, and the user +// confirms the code on a page that shows which account signs in, so a silent +// wrong-account answer is not the failure mode there. +type accountPromptForcer interface { + ForceAccountPrompt() +} + // AuthFlowInfo holds information for the OAuth 2.0 authorization flow type AuthFlowInfo struct { //nolint:revive DeviceCode string `json:"device_code"` @@ -49,6 +58,23 @@ type TokenInfo struct { ExpiresIn int `json:"expires_in"` UseIDToken bool `json:"-"` Email string `json:"-"` + EmailClaim string `json:"-"` +} + +// MatchesAccount reports whether the token belongs to the account a profile is +// bound to. A hint the IdP could not have acted on — no hint stored, or a token +// that carried no email claim — is reported as a match: the check exists to catch a +// login answered from the wrong account, not to block one it cannot judge. +// +// The comparison is case-insensitive. Local-parts are case-sensitive per RFC +// 5321, but no IdP in practice issues two accounts differing only in case, and +// an IdP that echoes a differently-cased address would otherwise fail every +// login. +func (t TokenInfo) MatchesAccount(hint string) bool { + if hint == "" || t.EmailClaim == "" { + return true + } + return strings.EqualFold(t.EmailClaim, hint) } // GetTokenToUse returns either the access or id token based on UseIDToken field @@ -63,19 +89,22 @@ func shouldUseDeviceFlow(force bool, isUnixDesktopClient bool) bool { return force || (runtime.GOOS == "linux" || runtime.GOOS == "freebsd") && !isUnixDesktopClient } -// NewOAuthFlow initializes and returns the appropriate OAuth flow based on the management configuration +// NewOAuthFlow initializes and returns the appropriate OAuth flow based on the management configuration. // // It starts by initializing the PKCE.If this process fails, it resorts to the Device Code Flow, // and if that also fails, the authentication process is deemed unsuccessful // -// On Linux distros without desktop environment support, it only tries to initialize the Device Code Flow -// forceDeviceCodeFlow can be used to skip PKCE and go directly to Device Code Flow (e.g., for Android TV) -func NewOAuthFlow(ctx context.Context, config *profilemanager.Config, isUnixDesktopClient bool, forceDeviceCodeFlow bool, hint string) (OAuthFlow, error) { +// On Linux distros without desktop environment support, it only tries to initialize the Device Code Flow. +// forceDeviceCodeFlow can be used to skip PKCE and go directly to Device Code Flow (e.g., for Android TV). +// +// sessionExtend marks the flow as renewing an existing peer's session rather than +// logging one in. See PKCEAuthorizationFlowRequest for what the server makes of it. +func NewOAuthFlow(ctx context.Context, config *profilemanager.Config, isUnixDesktopClient bool, forceDeviceCodeFlow bool, hint string, sessionExtend bool) (OAuthFlow, error) { if shouldUseDeviceFlow(forceDeviceCodeFlow, isUnixDesktopClient) { return authenticateWithDeviceCodeFlow(ctx, config, hint) } - pkceFlow, err := authenticateWithPKCEFlow(ctx, config, hint) + pkceFlow, err := authenticateWithPKCEFlow(ctx, config, hint, sessionExtend) if err != nil { log.Debugf("failed to initialize pkce authentication with error: %v\n", err) log.Debug("falling back to device code flow") @@ -85,14 +114,14 @@ func NewOAuthFlow(ctx context.Context, config *profilemanager.Config, isUnixDesk } // authenticateWithPKCEFlow initializes the Proof Key for Code Exchange flow auth flow -func authenticateWithPKCEFlow(ctx context.Context, config *profilemanager.Config, hint string) (OAuthFlow, error) { +func authenticateWithPKCEFlow(ctx context.Context, config *profilemanager.Config, hint string, sessionExtend bool) (OAuthFlow, error) { authClient, err := NewAuth(ctx, config.PrivateKey, config.ManagementURL, config) if err != nil { return nil, fmt.Errorf("failed to create auth client: %v", err) } defer authClient.Close() - pkceFlowInfo, err := authClient.getPKCEFlow(authClient.client) + pkceFlowInfo, err := authClient.getPKCEFlow(authClient.client, sessionExtend) if err != nil { return nil, fmt.Errorf("getting pkce authorization flow info failed with error: %v", err) } @@ -129,3 +158,21 @@ func authenticateWithDeviceCodeFlow(ctx context.Context, config *profilemanager. return deviceFlowInfo, nil } + +// RetryFlowForAccount returns a flow that asks the IdP to re-authenticate, for +// a login answered with an account other than the one hinted. Returns nil when +// the flow cannot ask — the caller then proceeds with the token it has. +// +// Proceeding rather than failing is deliberate. The hint is an email that may +// simply have changed since it was stored, and refusing the login would lock a +// user out of their own profile over a rename. The retry gives the account a +// chance to be corrected; the server still rejects a token that does not own +// the peer. +func RetryFlowForAccount(flow OAuthFlow) OAuthFlow { + forcer, ok := flow.(accountPromptForcer) + if !ok { + return nil + } + forcer.ForceAccountPrompt() + return flow +} diff --git a/client/internal/auth/pkce_flow.go b/client/internal/auth/pkce_flow.go index be64cc6a8..ebb0d4f1f 100644 --- a/client/internal/auth/pkce_flow.go +++ b/client/internal/auth/pkce_flow.go @@ -87,10 +87,11 @@ func validatePKCEConfig(config *PKCEAuthProviderConfig) error { // PKCEAuthorizationFlow implements the OAuthFlow interface for // the Authorization Code Flow with PKCE. type PKCEAuthorizationFlow struct { - providerConfig PKCEAuthProviderConfig - state string - codeVerifier string - oAuthConfig *oauth2.Config + providerConfig PKCEAuthProviderConfig + state string + codeVerifier string + oAuthConfig *oauth2.Config + forceAccountPrompt bool } // NewPKCEAuthorizationFlow returns new PKCE authorization code flow. @@ -153,11 +154,16 @@ func (p *PKCEAuthorizationFlow) RequestAuthInfo(ctx context.Context) (AuthFlowIn oauth2.SetAuthURLParam("code_challenge", codeChallenge), oauth2.SetAuthURLParam("audience", p.providerConfig.Audience), } + forceAccountPrompt := p.forceAccountPrompt + p.forceAccountPrompt = false + if !p.providerConfig.DisablePromptLogin { - switch p.providerConfig.LoginFlag { - case common.LoginFlagPromptLogin: + switch { + case forceAccountPrompt: params = append(params, oauth2.SetAuthURLParam("prompt", "login")) - case common.LoginFlagMaxAge0: + case p.providerConfig.LoginFlag == common.LoginFlagPromptLogin: + params = append(params, oauth2.SetAuthURLParam("prompt", "login")) + case p.providerConfig.LoginFlag == common.LoginFlagMaxAge0: params = append(params, oauth2.SetAuthURLParam("max_age", "0")) } } @@ -178,6 +184,20 @@ func (p *PKCEAuthorizationFlow) SetLoginHint(hint string) { p.providerConfig.LoginHint = hint } +// ForceAccountPrompt makes the next authorization request ask the IdP to +// re-authenticate instead of answering from the session it already holds. Used +// to retry a login that came back for an account other than the one hinted. +// +// The next RequestAuthInfo consumes the flag, so a flow that outlives its retry +// goes back to the configured behaviour instead of re-authenticating forever. +// +// DisablePromptLogin still wins: it is set for IdPs that break on prompt=login, +// where retrying with it would replace a wrong-account login with one that +// cannot complete at all. +func (p *PKCEAuthorizationFlow) ForceAccountPrompt() { + p.forceAccountPrompt = true +} + // WaitToken waits for the OAuth token in the PKCE Authorization Flow. // It starts an HTTP server to receive the OAuth token callback and waits for the token or an error. // Once the token is received, it is converted to TokenInfo and validated before returning. @@ -310,49 +330,48 @@ func (p *PKCEAuthorizationFlow) parseOAuthToken(token *oauth2.Token) (TokenInfo, return TokenInfo{}, fmt.Errorf("authentication failed: invalid access token - %w", err) } - email, err := parseEmailFromIDToken(tokenInfo.IDToken) + email, fromEmailClaim, err := parseEmailFromIDToken(tokenInfo.IDToken) if err != nil { log.Warnf("failed to parse email from ID token: %v", err) } else { tokenInfo.Email = email + if fromEmailClaim { + tokenInfo.EmailClaim = email + } } return tokenInfo, nil } // parseEmailFromIDToken extracts the email (or name) claim from an ID token -// without verifying its signature. The value is best-effort and used only as a -// UX convenience (login hint prefill and display); it never drives an -// authorization decision. The authoritative identity is established server-side -// from the signature-verified token. -func parseEmailFromIDToken(token string) (string, error) { +// without verifying its signature. The value is best-effort: it prefills the +// login hint and is displayed. Account matching (see MatchesAccount) only uses +// it when it came from the email claim, which fromEmailClaim reports. It never +// grants anything — the authoritative identity is established server-side from +// the signature-verified token. +func parseEmailFromIDToken(token string) (value string, fromEmailClaim bool, err error) { parts := strings.Split(token, ".") if len(parts) < 2 { - return "", fmt.Errorf("invalid token format") + return "", false, fmt.Errorf("invalid token format") } data, err := base64.RawURLEncoding.DecodeString(parts[1]) if err != nil { - return "", fmt.Errorf("failed to decode payload: %w", err) + return "", false, fmt.Errorf("failed to decode payload: %w", err) } var claims map[string]interface{} if err := json.Unmarshal(data, &claims); err != nil { - return "", fmt.Errorf("json unmarshal error: %w", err) + return "", false, fmt.Errorf("json unmarshal error: %w", err) } - var email string - if emailValue, ok := claims["email"].(string); ok { - email = emailValue - } else { - val, ok := claims["name"].(string) - if ok { - email = val - } else { - return "", fmt.Errorf("email or name field not found in token payload") - } + if email, ok := claims["email"].(string); ok { + return email, true, nil + } + if name, ok := claims["name"].(string); ok { + return name, false, nil } - return email, nil + return "", false, fmt.Errorf("email or name field not found in token payload") } func createCodeChallenge(codeVerifier string) string { diff --git a/client/internal/auth/pkce_flow_test.go b/client/internal/auth/pkce_flow_test.go index c487c13df..ccbca10e9 100644 --- a/client/internal/auth/pkce_flow_test.go +++ b/client/internal/auth/pkce_flow_test.go @@ -76,6 +76,32 @@ func TestPromptLogin(t *testing.T) { } } +func TestForceAccountPromptAppliesOnlyToTheRetry(t *testing.T) { + config := PKCEAuthProviderConfig{ + ClientID: "test-client-id", + Audience: "test-audience", + TokenEndpoint: "https://test-token-endpoint.com/token", + Scope: "openid email profile", + AuthorizationEndpoint: "https://test-auth-endpoint.com/authorize", + RedirectURLs: []string{"http://127.0.0.1:33992/"}, + UseIDToken: true, + LoginFlag: mgm.LoginFlagNone, + } + pkce, err := NewPKCEAuthorizationFlow(config) + require.NoError(t, err) + + pkce.ForceAccountPrompt() + + retry, err := pkce.RequestAuthInfo(context.Background()) + require.NoError(t, err) + require.Contains(t, retry.VerificationURIComplete, "prompt=login") + + next, err := pkce.RequestAuthInfo(context.Background()) + require.NoError(t, err) + require.NotContains(t, next.VerificationURIComplete, "prompt=login", + "the forced prompt outlived the retry it was armed for") +} + func TestIsPortInExcludedRange(t *testing.T) { tests := []struct { name string diff --git a/client/internal/auth/sessionwatch/watcher.go b/client/internal/auth/sessionwatch/watcher.go index e685c28d0..f71689df1 100644 --- a/client/internal/auth/sessionwatch/watcher.go +++ b/client/internal/auth/sessionwatch/watcher.go @@ -4,12 +4,20 @@ // T-WarningLead notification and a dismiss-gated T-FinalWarningLead // fallback dialog. // +// The deadline is an absolute wall-clock instant, so the watcher compares +// it against the wall clock on a ticker rather than arming a relative +// timer for it. A relative timer runs on the monotonic clock, which does +// not advance while a device is suspended: it fires once that much awake +// time has passed, which can be long after the deadline, and nothing +// re-evaluates when the device wakes up. Polling makes every tick after a +// resume (or after an NTP correction) see the real remaining time. +// // The watcher is idempotent: Update may be called as often as the network // map snapshots arrive. Repeating the same deadline is a no-op; a new -// deadline reschedules the timers and arms a fresh warning cycle. +// deadline starts a fresh warning cycle. // -// Warning firing is edge-detected. Each unique deadline value fires each -// warning callback at most once. +// Warning firing is edge-detected. Each unique deadline value publishes +// each warning at most once. package sessionwatch import ( @@ -28,9 +36,15 @@ const ( // maxDeadlineHorizon caps how far in the future an accepted deadline // can sit. A timestamp beyond this is almost certainly a protocol - // glitch, and silently arming a 100-year timer would hide the bug. + // glitch, and silently tracking a 100-year deadline would hide the bug. maxDeadlineHorizon = 10 * 365 * 24 * time.Hour + // defaultEvalInterval is how often the tracked deadline is compared + // against the wall clock. The leads are minutes, so a coarse tick costs + // nothing in accuracy and keeps the wakeup cheap on battery-powered + // devices. + defaultEvalInterval = 10 * time.Second + // WarningLead is how far before expiry the first (interactive) // warning fires. Drives the T-10 OS notification with // Extend/Dismiss actions. @@ -90,18 +104,23 @@ type StatusRecorder interface { // fallback T-FinalWarningLead dialog (suppressed when the user dismissed // the first one for the same deadline). Safe for concurrent use. type Watcher struct { - lead time.Duration - finalLead time.Duration + lead time.Duration + finalLead time.Duration + interval time.Duration + deadlineOnly bool // record the deadline, leave the warnings to the caller mu sync.Mutex current time.Time - timer *time.Timer - finalTimer *time.Timer - firedAt time.Time // deadline value the T-WarningLead callback last fired against - finalFiredAt time.Time // deadline value the T-FinalWarningLead callback last fired against - dismissedAt time.Time // deadline value the user dismissed via Dismiss(); gates fireFinal + firedAt time.Time // deadline value the T-WarningLead warning last published for + finalFiredAt time.Time // deadline value the T-FinalWarningLead warning last published for + dismissedAt time.Time // deadline value the user dismissed via Dismiss(); gates the final warning + announcedAt time.Time // deadline value the recorder has been told about; gates publishing closed bool recorder StatusRecorder + nowFn func() time.Time + stop chan struct{} // closed to stop the evaluation loop; nil while it is not running + done chan struct{} // closed by the loop on its way out + wake chan struct{} // buffered nudge asking the loop to evaluate before its next tick } // New returns a watcher with the package defaults WarningLead and @@ -113,27 +132,40 @@ func New(recorder StatusRecorder) *Watcher { } // NewWithLeads returns a watcher with custom lead times. Useful for tests. -// final must be strictly less than lead; otherwise both timers fire in the -// wrong order or simultaneously and the UI flow breaks. A zero final lead -// disables the final-warning timer entirely (see armTimerLocked) so a -// millisecond-scale deadline doesn't flush both timers in one tick. +// final must be strictly less than lead; otherwise the final warning takes +// over the whole warning window and the interactive notification never +// shows. A zero final lead disables the final warning entirely (see +// evaluate), leaving the interactive one as the only warning. func NewWithLeads(lead, final time.Duration, recorder StatusRecorder) *Watcher { return &Watcher{ lead: lead, finalLead: final, + interval: defaultEvalInterval, recorder: recorder, + nowFn: time.Now, } } +// NewDeadlineOnly returns a watcher that validates and records deadlines +// but publishes no warnings about them, and runs no evaluation loop to +// decide. Used where the deadline is handed on to something that schedules +// the warnings itself, such as the Android app. +func NewDeadlineOnly(recorder StatusRecorder) *Watcher { + w := New(recorder) + w.deadlineOnly = true + return w +} + // Update sets the latest deadline. Pass the zero time to clear (e.g. when // a Sync push from the server omits the field because login expiration // was disabled). // -// Same-value updates are no-ops. A different non-zero value cancels any -// pending timer, resets the "already fired" guards, and — when the -// deadline lies in the future — arms fresh warning timers. A deadline -// already in the past (within maxPastHorizon) is recorded as-is with no -// timers: the session has expired and consumers render it that way. +// Same-value updates are no-ops. A different non-zero value resets the +// "already fired" guards and starts a fresh warning cycle, evaluated +// immediately so a deadline that already sits inside a warning window +// warns without waiting for the next tick. A deadline already in the past +// (within maxPastHorizon) is recorded as-is and warns nothing: the session +// has expired and consumers render it that way. // // Returns one of the sentinel Err* values when the deadline fails the // sanity checks (pre-epoch, far future, or past beyond maxPastHorizon). @@ -153,7 +185,7 @@ func (w *Watcher) Update(deadline time.Time) error { return nil } - now := time.Now() + now := w.nowFn() switch { case deadline.Before(time.Unix(0, 0)): w.clearLocked() @@ -171,18 +203,22 @@ func (w *Watcher) Update(deadline time.Time) error { return nil } - w.stopTimerLocked() w.current = deadline - // Reset every per-deadline guard so a refreshed deadline arms a fresh + // Reset every per-deadline guard so a refreshed deadline starts a fresh // warning cycle: both edge triggers and the user Dismiss decision // (the user agreed to the old deadline expiring; a new deadline // restarts the contract). w.firedAt = time.Time{} w.finalFiredAt = time.Time{} w.dismissedAt = time.Time{} + w.announcedAt = time.Time{} - if deadline.After(now) { - w.armTimerLocked(deadline) + // Poll every accepted deadline, including one that reads as already + // expired: the clock may be running ahead of real time and get + // corrected later, and evaluate ignores a deadline that has genuinely + // passed anyway. + if !w.deadlineOnly { + w.startPollLocked() } recorder := w.recorder w.mu.Unlock() @@ -190,6 +226,28 @@ func (w *Watcher) Update(deadline time.Time) error { recorder.SetSessionExpiresAt(deadline) } log.Infof("auth session deadline set to: %s (in %s)", deadline.Format(time.RFC3339), time.Until(deadline).Round(time.Second)) + + // Open the gate only once the recorder knows the new deadline, so a + // warning that refers to it can never reach consumers before the state + // change itself. A tick landing in between finds the gate shut. + w.mu.Lock() + if w.closed || !w.current.Equal(deadline) { + w.mu.Unlock() + return nil + } + w.announcedAt = deadline + wake := w.wake + w.mu.Unlock() + + // Hand the evaluation to the loop rather than running it here, so a + // deadline that already sits inside a warning window is published at + // once without this goroutine ever touching the recorder. + if wake != nil { + select { + case wake <- struct{}{}: + default: + } + } return nil } @@ -202,9 +260,9 @@ func (w *Watcher) Deadline() time.Time { } // Dismiss records the user's "Dismiss" action against the current deadline -// and suppresses the upcoming final-warning callback for that deadline. -// Idempotent: repeated calls are no-ops. A subsequent Update with a fresh -// deadline resets the dismissal so the final-warning cycle re-arms. +// and suppresses the final warning for that deadline. Idempotent: repeated +// calls are no-ops. A subsequent Update with a fresh deadline resets the +// dismissal so the final-warning cycle starts over. // // No-op when the watcher holds no deadline or has been closed. func (w *Watcher) Dismiss() { @@ -217,35 +275,48 @@ func (w *Watcher) Dismiss() { return } w.dismissedAt = w.current - // Cancel the armed final-warning timer eagerly. fireFinal would also - // gate on dismissedAt, but stopping the timer avoids a wakeup with - // nothing to do and makes the intent visible. - if w.finalTimer != nil { - w.finalTimer.Stop() - w.finalTimer = nil - } log.Infof("auth session final-warning dismissed for deadline %s", w.current.Format(time.RFC3339)) } -// Close stops any pending timer. Update calls after Close are ignored. -// The recorder keeps its deadline: the watcher is engine-scoped and closes -// on every engine restart (network change, sleep/wake, stream errors) -// while the SSO deadline stays valid across those, so clearing here would -// blank the UI's "expires in" row on every transient reconnect. The -// client run loop clears the server-scoped recorder when it exits for -// real (Down, profile switch, permanent login failure). +// Close stops the evaluation loop and waits for it to exit. Update calls +// after Close are ignored. The recorder keeps its deadline: the watcher is +// engine-scoped and closes on every engine restart (network change, +// sleep/wake, stream errors) while the SSO deadline stays valid across +// those, so clearing here would blank the UI's "expires in" row on every +// transient reconnect. The client run loop clears the server-scoped +// recorder when it exits for real (Down, profile switch, permanent login +// failure). func (w *Watcher) Close() { w.mu.Lock() - defer w.mu.Unlock() if w.closed { + // A concurrent Close is already tearing down. w.done outlives it, + // so this caller waits for the same loop rather than returning + // while a warning is still on its way to the recorder. + done := w.done + w.mu.Unlock() + if done != nil { + <-done + } return } w.closed = true - w.stopTimerLocked() w.current = time.Time{} w.firedAt = time.Time{} w.finalFiredAt = time.Time{} w.dismissedAt = time.Time{} + w.announcedAt = time.Time{} + // Copy the channels out before releasing the lock: the loop takes w.mu + // on every tick, so waiting for it while holding the lock would + // deadlock. w.done stays on the receiver for the branch above. + stop, done := w.stop, w.done + w.stop, w.wake = nil, nil + w.mu.Unlock() + + if stop == nil { + return + } + close(stop) + <-done } // clearLocked drops the tracked deadline and notifies the recorder so @@ -257,11 +328,11 @@ func (w *Watcher) clearLocked() { w.mu.Unlock() return } - w.stopTimerLocked() w.current = time.Time{} w.firedAt = time.Time{} w.finalFiredAt = time.Time{} w.dismissedAt = time.Time{} + w.announcedAt = time.Time{} recorder := w.recorder w.mu.Unlock() if recorder != nil { @@ -270,87 +341,121 @@ func (w *Watcher) clearLocked() { log.Infof("auth session deadline cleared") } -func (w *Watcher) stopTimerLocked() { - if w.timer != nil { - w.timer.Stop() - w.timer = nil +// startPollLocked starts the evaluation loop unless it is already +// running. The loop starts lazily on the first accepted deadline, so a +// client whose server never publishes a session expiry never pays for a +// ticker, and it runs until Close: the watcher is engine-scoped, and a +// cleared deadline is normally followed by a fresh one on the next sync. +// Caller must hold w.mu. +func (w *Watcher) startPollLocked() { + if w.stop != nil { + return } - if w.finalTimer != nil { - w.finalTimer.Stop() - w.finalTimer = nil + stop := make(chan struct{}) + done := make(chan struct{}) + wake := make(chan struct{}, 1) + w.stop, w.done, w.wake = stop, done, wake + go w.poll(stop, done, wake, w.interval) +} + +// poll re-evaluates the tracked deadline every interval, and as soon as a +// new deadline is announced. It is the only caller of evaluate, so Close +// waiting for it to exit is enough to know no warning is still on its way +// to the recorder. Its channels and interval are passed in rather than +// read off the receiver, so Close can clear them without racing it. +func (w *Watcher) poll(stop <-chan struct{}, done chan<- struct{}, wake <-chan struct{}, interval time.Duration) { + defer close(done) + + ticker := time.NewTicker(interval) + defer ticker.Stop() + + for { + select { + case <-stop: + return + case <-wake: + w.evaluate() + case <-ticker.C: + w.evaluate() + } } } -func (w *Watcher) armTimerLocked(deadline time.Time) { - w.timer = armOneShotLocked(deadline.Add(-w.lead), func() { w.fire(deadline) }) - // finalLead <= 0 disables the final-warning timer entirely. Used by - // tests that predate the final-warning fallback so a millisecond-scale - // deadline does not flush both timers at once. - if w.finalLead > 0 { - w.finalTimer = armOneShotLocked(deadline.Add(-w.finalLead), func() { w.fireFinal(deadline) }) - } -} - -func (w *Watcher) fire(armedFor time.Time) { +// evaluate compares the tracked deadline against the wall clock and +// publishes whichever warning the remaining time calls for. Inside the +// final-warning window the interactive warning is stale, so the final one +// is published in its place: a device that resumes there never saw the +// T-WarningLead notification. +func (w *Watcher) evaluate() { w.mu.Lock() - if w.closed || !w.current.Equal(armedFor) { - // Deadline moved while we were waiting (e.g. a successful extend). - // The reschedule path armed a fresh timer; this one is stale. + if w.closed || w.deadlineOnly || w.current.IsZero() { w.mu.Unlock() return } - if !w.firedAt.IsZero() && w.firedAt.Equal(armedFor) { + + deadline := w.current + if !w.announcedAt.Equal(deadline) { + // Update is still on its way to the recorder with this deadline. w.mu.Unlock() return } - w.firedAt = armedFor + // Round(0) strips the monotonic reading so the comparison is wall + // clock on both sides, whether the deadline came off the wire or from + // a caller that derived it from time.Now. + remaining := deadline.Round(0).Sub(w.nowFn().Round(0)) + + switch { + case remaining <= 0: + // Already expired: the post-mortem SessionExpired flow owns it. + w.mu.Unlock() + case w.finalLead > 0 && remaining <= w.finalLead: + w.publishFinalLocked(deadline, remaining) + case remaining <= w.lead: + w.publishWarningLocked(deadline, remaining) + default: + w.mu.Unlock() + } +} + +// publishWarningLocked emits the interactive T-WarningLead warning, at +// most once per deadline value. Caller must hold w.mu; this helper +// releases it. +func (w *Watcher) publishWarningLocked(deadline time.Time, remaining time.Duration) { + if w.firedAt.Equal(deadline) { + w.mu.Unlock() + return + } + w.firedAt = deadline recorder := w.recorder w.mu.Unlock() if recorder == nil { return } - log.Infof("auth session expiry soon warning fired") - publishWarning(recorder, armedFor, false) + log.Infof("auth session expiry soon warning fired for deadline %s (in %s)", + deadline.Format(time.RFC3339), remaining.Round(time.Second)) + publishWarning(recorder, deadline, false) } -// fireFinal mirrors fire for the T-FinalWarningLead timer with an extra -// dismiss-gate: if the user dismissed the T-WarningLead notification for -// this deadline, the final warning is suppressed entirely. -func (w *Watcher) fireFinal(armedFor time.Time) { - w.mu.Lock() - if w.closed || !w.current.Equal(armedFor) { +// publishFinalLocked emits the final warning, at most once per deadline +// value and never once the user dismissed that deadline. It marks the +// interactive warning as handled too: the final window is open, so a +// "expires in WarningLead minutes" notification would be wrong. Caller +// must hold w.mu; this helper releases it. +func (w *Watcher) publishFinalLocked(deadline time.Time, remaining time.Duration) { + if w.finalFiredAt.Equal(deadline) || w.dismissedAt.Equal(deadline) { w.mu.Unlock() return } - if !w.finalFiredAt.IsZero() && w.finalFiredAt.Equal(armedFor) { - w.mu.Unlock() - return - } - if w.dismissedAt.Equal(armedFor) { - w.mu.Unlock() - log.Infof("auth session final-warning skipped (dismissed by user)") - return - } - w.finalFiredAt = armedFor + w.firedAt = deadline + w.finalFiredAt = deadline recorder := w.recorder w.mu.Unlock() if recorder == nil { return } - log.Infof("auth session final-warning fired") - publishWarning(recorder, armedFor, true) -} - -// armOneShotLocked schedules cb at fireAt. When fireAt is already in the -// past it dispatches on the next scheduler tick so a state-change recorder -// notification (invoked after w.mu is released) lands first. Caller must -// hold w.mu. -func armOneShotLocked(fireAt time.Time, cb func()) *time.Timer { - delay := time.Until(fireAt) - if delay <= 0 { - return time.AfterFunc(0, cb) - } - return time.AfterFunc(delay, cb) + log.Infof("auth session final-warning fired for deadline %s (in %s)", + deadline.Format(time.RFC3339), remaining.Round(time.Second)) + publishWarning(recorder, deadline, true) } // publishWarning composes the SystemEvent for a watcher-fired warning and diff --git a/client/internal/auth/sessionwatch/watcher_test.go b/client/internal/auth/sessionwatch/watcher_test.go index 4b49a94b6..f0cc3ed88 100644 --- a/client/internal/auth/sessionwatch/watcher_test.go +++ b/client/internal/auth/sessionwatch/watcher_test.go @@ -19,6 +19,9 @@ type fakeRecorder struct { mu sync.Mutex events []event lastDeadline time.Time + // setDelay stalls SetSessionExpiresAt, widening the window in which a + // concurrent evaluation could publish a warning out of order. + setDelay time.Duration } type eventKind int @@ -43,6 +46,12 @@ type event struct { // is the zero time, so an initial clear before any deadline is set emits // nothing — matching the real recorder. func (r *fakeRecorder) SetSessionExpiresAt(deadline time.Time) { + r.mu.Lock() + delay := r.setDelay + r.mu.Unlock() + // Stall without the lock held, so a concurrent publish can still record. + time.Sleep(delay) + r.mu.Lock() defer r.mu.Unlock() if r.lastDeadline.Equal(deadline) { @@ -116,10 +125,54 @@ func waitForEvents(t *testing.T, r *fakeRecorder, want int) []event { return nil } -// newWatcher builds a watcher with the final timer disabled (finalLead=0), +// testInterval keeps the ticker-driven tests fast; the leads they use are +// in the same millisecond scale. +const testInterval = 2 * time.Millisecond + +// newWatcher builds a watcher with the final warning disabled (finalLead=0), // matching the lead-only behaviour the pre-final-warning tests assume. func newWatcher(lead time.Duration, r *fakeRecorder) *Watcher { - return NewWithLeads(lead, 0, r) + return newWatcherWithLeads(lead, 0, r) +} + +// newWatcherWithLeads builds a watcher that evaluates on testInterval. The +// interval is set before Update, so the evaluation loop does not exist yet +// and the write cannot race it. +func newWatcherWithLeads(lead, final time.Duration, r StatusRecorder) *Watcher { + w := NewWithLeads(lead, final, r) + w.interval = testInterval + return w +} + +// fakeClock is the watcher's wall clock under test. Reads come from the +// evaluation goroutine while the test writes, so both go through the mutex. +type fakeClock struct { + mu sync.Mutex + t time.Time +} + +func newFakeClock(t time.Time) *fakeClock { + return &fakeClock{t: t} +} + +func (c *fakeClock) now() time.Time { + c.mu.Lock() + defer c.mu.Unlock() + return c.t +} + +// set jumps the clock, standing in for a resume from suspension or for an +// NTP correction. +func (c *fakeClock) set(t time.Time) { + c.mu.Lock() + defer c.mu.Unlock() + c.t = t +} + +// settle waits out a handful of evaluation ticks so a test can assert that +// nothing was published. +func settle() { + time.Sleep(20 * testInterval) } func TestUpdateZeroBeforeAnythingIsNoop(t *testing.T) { @@ -410,7 +463,7 @@ func TestCloseWithoutDeadlineLeavesRecorderUntouched(t *testing.T) { func TestFinalWarningFiresAfterRegularWarning(t *testing.T) { r := &fakeRecorder{} // Warning fires at deadline-80ms, final at deadline-30ms. - w := NewWithLeads(80*time.Millisecond, 30*time.Millisecond, r) + w := newWatcherWithLeads(80*time.Millisecond, 30*time.Millisecond, r) defer w.Close() d := time.Now().Add(100 * time.Millisecond) @@ -446,7 +499,7 @@ func TestFinalWarningFiresAfterRegularWarning(t *testing.T) { func TestDismissSuppressesFinalWarning(t *testing.T) { r := &fakeRecorder{} - w := NewWithLeads(80*time.Millisecond, 30*time.Millisecond, r) + w := newWatcherWithLeads(80*time.Millisecond, 30*time.Millisecond, r) defer w.Close() d := time.Now().Add(100 * time.Millisecond) @@ -477,7 +530,7 @@ func TestDismissSuppressesFinalWarning(t *testing.T) { func TestDismissResetByNewDeadline(t *testing.T) { r := &fakeRecorder{} - w := NewWithLeads(80*time.Millisecond, 30*time.Millisecond, r) + w := newWatcherWithLeads(80*time.Millisecond, 30*time.Millisecond, r) defer w.Close() first := time.Now().Add(100 * time.Millisecond) @@ -507,7 +560,7 @@ func TestDismissResetByNewDeadline(t *testing.T) { func TestDismissBeforeUpdateIsNoop(t *testing.T) { r := &fakeRecorder{} - w := NewWithLeads(80*time.Millisecond, 30*time.Millisecond, r) + w := newWatcherWithLeads(80*time.Millisecond, 30*time.Millisecond, r) defer w.Close() // No deadline tracked yet; Dismiss must be a no-op (no panic, no state). @@ -527,3 +580,340 @@ func TestDismissBeforeUpdateIsNoop(t *testing.T) { } t.Fatalf("final-warning did not publish after no-op pre-Update Dismiss, events=%+v", r.snapshot()) } + +// The tests below drive the watcher's wall clock directly. The deadline sits +// an hour out in real time and the fake clock jumps, standing in for a device +// that was suspended and resumed somewhere inside — or past — the warning +// windows. The evaluation loop is what reacts to the jump, so these exercise +// the same path production takes on a resume. +func newResumeWatcher(t *testing.T, r *fakeRecorder) (*Watcher, *fakeClock, time.Time) { + t.Helper() + + w := newWatcherWithLeads(WarningLead, FinalWarningLead, r) + start := time.Now() + clock := newFakeClock(start) + // Set before Update: the evaluation loop does not exist yet. + w.nowFn = clock.now + t.Cleanup(w.Close) + + deadline := start.Add(time.Hour).Round(0) + if err := w.Update(deadline); err != nil { + t.Fatalf("Update: %v", err) + } + return w, clock, deadline +} + +func TestResumeInsideWarningWindowWarns(t *testing.T) { + r := &fakeRecorder{} + _, clock, deadline := newResumeWatcher(t, r) + + clock.set(deadline.Add(-5 * time.Minute)) + + events := waitForEvents(t, r, 2) + if !events[1].isWarning() { + t.Fatalf("expected the interactive warning after the resume, got %+v", events[1]) + } + if n := countWhere(events, event.isFinalWarning); n != 0 { + t.Fatalf("final-warning must wait for its own window, got %d: %+v", n, events) + } +} + +func TestResumeInsideFinalWindowSendsFinalWarningOnly(t *testing.T) { + r := &fakeRecorder{} + _, clock, deadline := newResumeWatcher(t, r) + + clock.set(deadline.Add(-time.Minute)) + + events := waitForEvents(t, r, 2) + if !events[1].isFinalWarning() { + t.Fatalf("expected the final warning after the resume, got %+v", events[1]) + } + settle() + if n := countWhere(r.snapshot(), event.isWarning); n != 0 { + t.Fatalf("the interactive warning is stale inside the final window, got %d: %+v", n, r.snapshot()) + } +} + +func TestResumePastDeadlinePublishesNothing(t *testing.T) { + r := &fakeRecorder{} + _, clock, deadline := newResumeWatcher(t, r) + + clock.set(deadline.Add(time.Minute)) + + settle() + if n := countWhere(r.snapshot(), func(e event) bool { return e.kind == publish }); n != 0 { + t.Fatalf("an expired session must not warn, got %d publishes: %+v", n, r.snapshot()) + } +} + +func TestWarningPublishesOncePerDeadline(t *testing.T) { + r := &fakeRecorder{} + _, clock, deadline := newResumeWatcher(t, r) + + clock.set(deadline.Add(-5 * time.Minute)) + waitForEvents(t, r, 2) + + // Many ticks pass inside the same window. + settle() + if n := countWhere(r.snapshot(), event.isWarning); n != 1 { + t.Fatalf("expected exactly 1 warning publish across ticks, got %d: %+v", n, r.snapshot()) + } +} + +// TestWarningRecoversFromAClockRunningAhead covers the device that boots +// before NTP has corrected it: the deadline looks long gone, nothing is +// published, and the warning still arrives once the clock is fixed. +func TestWarningRecoversFromAClockRunningAhead(t *testing.T) { + r := &fakeRecorder{} + _, clock, deadline := newResumeWatcher(t, r) + + clock.set(deadline.Add(2 * time.Hour)) + settle() + if n := countWhere(r.snapshot(), func(e event) bool { return e.kind == publish }); n != 0 { + t.Fatalf("a deadline that looks expired must not warn, got %d publishes: %+v", n, r.snapshot()) + } + + clock.set(deadline.Add(-5 * time.Minute)) + + events := waitForEvents(t, r, 2) + if !events[1].isWarning() { + t.Fatalf("expected the warning once the clock was corrected, got %+v", events[1]) + } +} + +func TestResumeInsideFinalWindowRespectsDismiss(t *testing.T) { + r := &fakeRecorder{} + w, clock, deadline := newResumeWatcher(t, r) + + // The user dismissed the warning for this deadline before the device + // was suspended; resuming inside the final window must not reopen it. + w.Dismiss() + clock.set(deadline.Add(-time.Minute)) + + settle() + if n := countWhere(r.snapshot(), func(e event) bool { return e.kind == publish }); n != 0 { + t.Fatalf("a dismissed deadline must not warn on resume, got %d: %+v", n, r.snapshot()) + } +} + +func TestCloseStopsTheEvaluationLoop(t *testing.T) { + r := &fakeRecorder{} + w := newWatcherWithLeads(WarningLead, FinalWarningLead, r) + start := time.Now() + clock := newFakeClock(start) + w.nowFn = clock.now + + deadline := start.Add(time.Hour).Round(0) + if err := w.Update(deadline); err != nil { + t.Fatalf("Update: %v", err) + } + w.Close() + + clock.set(deadline.Add(-5 * time.Minute)) + settle() + + if n := countWhere(r.snapshot(), func(e event) bool { return e.kind == publish }); n != 0 { + t.Fatalf("a closed watcher must not publish, got %d: %+v", n, r.snapshot()) + } +} + +func TestDeadlineOnlyRecordsDeadlineWithoutWarnings(t *testing.T) { + r := &fakeRecorder{} + w := NewDeadlineOnly(r) + defer w.Close() + + // With the default leads this deadline sits inside the final-warning + // window, so a watcher that warns at all would publish on the spot. + d := time.Now().Add(50 * time.Millisecond).Round(0) + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + if got := r.deadline(); !got.Equal(d) { + t.Fatalf("expected recorder deadline %v, got %v", d, got) + } + + time.Sleep(100 * time.Millisecond) + + events := r.snapshot() + if got := countWhere(events, func(e event) bool { return e.kind == publish }); got != 0 { + t.Fatalf("expected no publish in deadline-only mode, got %d: %+v", got, events) + } + w.mu.Lock() + polling := w.stop != nil + w.mu.Unlock() + if polling { + t.Fatal("expected no evaluation loop in deadline-only mode") + } +} + +func TestDeadlineOnlyStillRejectsOutOfRangeDeadlines(t *testing.T) { + r := &fakeRecorder{} + w := NewDeadlineOnly(r) + defer w.Close() + + if err := w.Update(time.Now().Add(time.Hour)); err != nil { + t.Fatalf("Update: %v", err) + } + + err := w.Update(time.Now().Add(-maxPastHorizon - time.Hour)) + if !errors.Is(err, ErrDeadlineInPast) { + t.Fatalf("expected ErrDeadlineInPast, got %v", err) + } + if got := r.deadline(); !got.IsZero() { + t.Fatalf("expected recorder cleared after rejection, got %v", got) + } +} + +// TestClockAheadAtUpdateStillWarnsOnceCorrected covers a client that connects +// while its clock runs ahead of real time: the deadline reads as already +// expired when it arrives, so nothing publishes, and the warning must still +// come once NTP pulls the clock back. +func TestClockAheadAtUpdateStillWarnsOnceCorrected(t *testing.T) { + r := &fakeRecorder{} + w := newWatcherWithLeads(WarningLead, FinalWarningLead, r) + deadline := time.Now().Add(time.Hour).Round(0) + // Two hours past the deadline from the device's point of view, well + // inside maxPastHorizon, so Update accepts and records it. + clock := newFakeClock(deadline.Add(2 * time.Hour)) + w.nowFn = clock.now + t.Cleanup(w.Close) + + if err := w.Update(deadline); err != nil { + t.Fatalf("Update: %v", err) + } + + settle() + if n := countWhere(r.snapshot(), func(e event) bool { return e.kind == publish }); n != 0 { + t.Fatalf("a deadline that reads as expired must not warn, got %d: %+v", n, r.snapshot()) + } + + clock.set(deadline.Add(-5 * time.Minute)) + + events := waitForEvents(t, r, 2) + if !events[1].isWarning() { + t.Fatalf("expected the warning once the clock was corrected, got %+v", events[1]) + } +} + +// TestWarningNeverPrecedesTheDeadlineStateChange pins the ordering Update +// documents: consumers learn the new deadline before they see a warning that +// refers to it. The deadline here already sits inside the warning window, and +// the recorder stalls, so the evaluation loop gets many chances to publish +// while Update is still announcing. +func TestWarningNeverPrecedesTheDeadlineStateChange(t *testing.T) { + r := &fakeRecorder{setDelay: 50 * time.Millisecond} + w := newWatcherWithLeads(WarningLead, FinalWarningLead, r) + deadline := time.Now().Add(time.Hour).Round(0) + clock := newFakeClock(deadline.Add(-5 * time.Minute)) + w.nowFn = clock.now + t.Cleanup(w.Close) + + if err := w.Update(deadline); err != nil { + t.Fatalf("Update: %v", err) + } + + events := waitForEvents(t, r, 2) + if events[0].kind != stateChange { + t.Fatalf("event[0] should be the deadline state change, got %+v", events) + } + if !events[1].isWarning() { + t.Fatalf("event[1] should be the warning, got %+v", events) + } +} + +// TestUpdateWakesTheLoopWithoutWaitingForATick pins the hand-off: Update +// publishes nothing itself, it nudges the evaluation loop, so a deadline that +// already sits inside a warning window is warned about at once even though the +// ticker here would not fire for an hour. +func TestUpdateWakesTheLoopWithoutWaitingForATick(t *testing.T) { + r := &fakeRecorder{} + w := newWatcherWithLeads(WarningLead, FinalWarningLead, r) + w.interval = time.Hour + t.Cleanup(w.Close) + + d := time.Now().Add(5 * time.Minute).Round(0) + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + + events := waitForEvents(t, r, 2) + if !events[1].isWarning() { + t.Fatalf("expected the warning on the wake-up, got %+v", events[1]) + } +} + +// blockingRecorder holds a publish open until the test releases it, so the +// evaluation loop can be parked mid-publish while Close is called. +type blockingRecorder struct { + fakeRecorder + entered chan struct{} + release chan struct{} +} + +func (r *blockingRecorder) PublishEvent( + severity cProto.SystemEvent_Severity, + category cProto.SystemEvent_Category, + message string, + userMessage string, + metadata map[string]string, +) { + select { + case r.entered <- struct{}{}: + default: + } + <-r.release + r.fakeRecorder.PublishEvent(severity, category, message, userMessage, metadata) +} + +// TestConcurrentCloseWaitsForTheLoop pins the contract for the caller that +// loses the race: Close returns only once the loop is done, whichever of the +// two calls got there first. +func TestConcurrentCloseWaitsForTheLoop(t *testing.T) { + r := &blockingRecorder{entered: make(chan struct{}, 1), release: make(chan struct{})} + var releaseOnce sync.Once + release := func() { releaseOnce.Do(func() { close(r.release) }) } + + w := newWatcherWithLeads(WarningLead, FinalWarningLead, r) + w.interval = time.Hour + // Order matters: Close waits for the loop, which stays parked in the + // publish until the release, so a t.Fatal below would hang the cleanup + // instead of reporting the failure. + t.Cleanup(func() { + release() + w.Close() + }) + + d := time.Now().Add(5 * time.Minute).Round(0) + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + + // The loop is now parked inside the warning publish. + select { + case <-r.entered: + case <-time.After(2 * time.Second): + t.Fatal("the loop never reached the publish") + } + + first := make(chan struct{}) + second := make(chan struct{}) + go func() { w.Close(); close(first) }() + go func() { w.Close(); close(second) }() + + select { + case <-first: + t.Fatal("Close returned while the loop was still publishing") + case <-second: + t.Fatal("Close returned while the loop was still publishing") + case <-time.After(100 * time.Millisecond): + } + + release() + for _, done := range []chan struct{}{first, second} { + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("Close did not return after the publish completed") + } + } +} diff --git a/client/internal/connect.go b/client/internal/connect.go index 88d829d2f..b42ef9818 100644 --- a/client/internal/connect.go +++ b/client/internal/connect.go @@ -292,6 +292,10 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan return nil } + if c.updateManager != nil { + c.updateManager.ResetMode() + } + // suspend connection attempts while the OS reports no usable network if waited, err := c.netMgr.Wait(c.ctx); err != nil { return nil diff --git a/client/internal/daemonaddr/grpc.go b/client/internal/daemonaddr/grpc.go new file mode 100644 index 000000000..5f0cc10cc --- /dev/null +++ b/client/internal/daemonaddr/grpc.go @@ -0,0 +1,42 @@ +package daemonaddr + +import ( + "os" + "strconv" + + log "github.com/sirupsen/logrus" +) + +const ( + // EnvMaxRecvMsgSize overrides the default gRPC max receive message size for + // connections to the daemon. Value is in bytes. + EnvMaxRecvMsgSize = "NB_DAEMON_GRPC_MAX_MSG_SIZE" + + // defaultMaxRecvMsgSize is the max gRPC receive message size used for daemon + // connections when EnvMaxRecvMsgSize is unset or invalid. It overrides the + // gRPC library default of 4 MB, which a detailed status already exceeds on a + // network of a few thousand peers. + defaultMaxRecvMsgSize = 1024 * 1024 * 16 +) + +// MaxRecvMsgSize returns the max gRPC receive message size for daemon connections +// from the environment, or defaultMaxRecvMsgSize (16 MB) if unset or invalid. +func MaxRecvMsgSize() int { + val := os.Getenv(EnvMaxRecvMsgSize) + if val == "" { + return defaultMaxRecvMsgSize + } + + size, err := strconv.Atoi(val) + if err != nil { + log.Warnf("invalid %s value %q, using default: %v", EnvMaxRecvMsgSize, val, err) + return defaultMaxRecvMsgSize + } + + if size <= 0 { + log.Warnf("invalid %s value %d, must be positive, using default", EnvMaxRecvMsgSize, size) + return defaultMaxRecvMsgSize + } + + return size +} diff --git a/client/internal/daemonaddr/grpc_test.go b/client/internal/daemonaddr/grpc_test.go new file mode 100644 index 000000000..7c4909a42 --- /dev/null +++ b/client/internal/daemonaddr/grpc_test.go @@ -0,0 +1,112 @@ +package daemonaddr + +import ( + "context" + "net" + "os" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/proto" +) + +func TestMaxRecvMsgSize(t *testing.T) { + tests := []struct { + name string + envValue string + expected int + }{ + {name: "unset returns default", envValue: "", expected: defaultMaxRecvMsgSize}, + {name: "non-numeric returns default", envValue: "abc", expected: defaultMaxRecvMsgSize}, + {name: "negative returns default", envValue: "-1", expected: defaultMaxRecvMsgSize}, + {name: "zero returns default", envValue: "0", expected: defaultMaxRecvMsgSize}, + {name: "valid value is used", envValue: "33554432", expected: 33554432}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + // Set first so the previous value is restored on cleanup, then unset to + // exercise the absent case. + t.Setenv(EnvMaxRecvMsgSize, tc.envValue) + if tc.envValue == "" { + require.NoError(t, os.Unsetenv(EnvMaxRecvMsgSize), "unset the override") + } + + assert.Equal(t, tc.expected, MaxRecvMsgSize(), "max receive message size") + }) + } +} + +// bigStatusServer answers Status with a response larger than gRPC's 4 MB default +// receive limit, which is what a detailed status on a large network looks like. +type bigStatusServer struct { + proto.UnimplementedDaemonServiceServer + payload string +} + +func (s *bigStatusServer) Status(context.Context, *proto.StatusRequest) (*proto.StatusResponse, error) { + return &proto.StatusResponse{Status: s.payload}, nil +} + +func startBigStatusServer(t *testing.T, payload string) string { + t.Helper() + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err, "listen on loopback") + + srv := grpc.NewServer() + proto.RegisterDaemonServiceServer(srv, &bigStatusServer{payload: payload}) + go func() { + _ = srv.Serve(listener) + }() + t.Cleanup(srv.Stop) + + return "tcp://" + listener.Addr().String() +} + +func TestDialTargetAcceptsAStatusOverTheGrpcDefault(t *testing.T) { + payload := strings.Repeat("x", 5*1024*1024) + addr := startBigStatusServer(t, payload) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + target, opts := DialTarget(addr) + conn, err := grpc.NewClient(target, opts...) + require.NoError(t, err, "dial the daemon") + t.Cleanup(func() { _ = conn.Close() }) + + resp, err := proto.NewDaemonServiceClient(conn).Status(ctx, &proto.StatusRequest{}) + require.NoError(t, err, "a detailed status must not be rejected for its size") + assert.Len(t, resp.GetStatus(), len(payload), "the whole response must arrive") +} + +// TestDialTargetRaisesTheDefaultLimit is the negative control: the same response +// over a connection carrying gRPC's own defaults is refused, which is the failure +// reported by `netbird status -d` on a large deployment. +func TestDialTargetRaisesTheDefaultLimit(t *testing.T) { + payload := strings.Repeat("x", 5*1024*1024) + addr := startBigStatusServer(t, payload) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + conn, err := grpc.NewClient( + strings.TrimPrefix(addr, "tcp://"), + grpc.WithTransportCredentials(insecure.NewCredentials()), + ) + require.NoError(t, err, "dial with the library defaults") + t.Cleanup(func() { _ = conn.Close() }) + + _, err = proto.NewDaemonServiceClient(conn).Status(ctx, &proto.StatusRequest{}) + require.Error(t, err, "the library default must reject this response") + assert.Equal(t, codes.ResourceExhausted, status.Code(err), "gRPC rejects an oversized message") +} diff --git a/client/internal/daemonaddr/pipe.go b/client/internal/daemonaddr/pipe.go index 51815ef5e..bf1c8fdd0 100644 --- a/client/internal/daemonaddr/pipe.go +++ b/client/internal/daemonaddr/pipe.go @@ -36,7 +36,10 @@ const ( // address. The npipe scheme needs a context dialer because gRPC has no // named-pipe resolver; unix and tcp are handled by gRPC itself. func DialTarget(addr string) (string, []grpc.DialOption) { - opts := []grpc.DialOption{grpc.WithTransportCredentials(insecure.NewCredentials())} + opts := []grpc.DialOption{ + grpc.WithTransportCredentials(insecure.NewCredentials()), + grpc.WithDefaultCallOptions(grpc.MaxCallRecvMsgSize(MaxRecvMsgSize())), + } if name, ok := strings.CutPrefix(addr, pipeScheme); ok { paths := PipePaths(name) diff --git a/client/internal/debug/debug.go b/client/internal/debug/debug.go index b362ae293..f4d1c3598 100644 --- a/client/internal/debug/debug.go +++ b/client/internal/debug/debug.go @@ -379,9 +379,38 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen } } +// bundleFilePattern names the bundle zips Generate creates in tempDir; the +// asterisk is filled in by os.CreateTemp. +const bundleFilePattern = "netbird.debug.*.zip" + +const exportedBundlePrefix = "netbird.debug-file." + +const exportedBundleMaxAge = 24 * time.Hour + +// RemoveStaleBundles deletes bundle zips that an interrupted generation or +// upload left behind in dir. Only files older than maxAge go, so a bundle that +// another caller is still writing or uploading in the same directory survives. +// Exported bundles are kept for exportedBundleMaxAge instead. +func RemoveStaleBundles(dir string, maxAge time.Duration) { + removeStaleFiles(dir, bundleFilePattern, maxAge) + removeStaleFiles(dir, exportedBundlePrefix+"*.zip", exportedBundleMaxAge) +} + +// ExportBundle renames a generated bundle out of the RemoveStaleBundles pattern +// and returns the new path. The caller owns the file from then on; an export +// abandoned for longer than exportedBundleMaxAge is removed by RemoveStaleBundles. +func ExportBundle(path string) (string, error) { + base := strings.TrimPrefix(filepath.Base(path), strings.SplitN(bundleFilePattern, "*", 2)[0]) + exported := filepath.Join(filepath.Dir(path), exportedBundlePrefix+base) + if err := os.Rename(path, exported); err != nil { + return "", fmt.Errorf("export debug bundle: %w", err) + } + return exported, nil +} + // Generate creates a debug bundle and returns the location. func (g *BundleGenerator) Generate() (resp string, err error) { - bundlePath, err := os.CreateTemp(g.tempDir, "netbird.debug.*.zip") + bundlePath, err := os.CreateTemp(g.tempDir, bundleFilePattern) if err != nil { return "", fmt.Errorf("create zip file: %w", err) } @@ -1725,3 +1754,26 @@ func anonymizeSlice(v []any, anonymizer *anonymize.Anonymizer) []any { } return v } + +func removeStaleFiles(dir, pattern string, maxAge time.Duration) { + matches, err := filepath.Glob(filepath.Join(dir, pattern)) + if err != nil { + log.Debugf("glob stale debug bundles in %s: %v", dir, err) + return + } + + cutoff := time.Now().Add(-maxAge) + for _, path := range matches { + info, err := os.Stat(path) + if err != nil || info.ModTime().After(cutoff) { + continue + } + if err := os.Remove(path); err != nil { + if !errors.Is(err, fs.ErrNotExist) { + log.Warnf("remove stale debug bundle %s: %v", path, err) + } + continue + } + log.Infof("removed stale debug bundle %s", path) + } +} diff --git a/client/internal/debug/debug_test.go b/client/internal/debug/debug_test.go index 17d520358..0f74490f2 100644 --- a/client/internal/debug/debug_test.go +++ b/client/internal/debug/debug_test.go @@ -4,6 +4,7 @@ import ( "archive/zip" "bytes" "encoding/json" + "fmt" "net" "net/netip" "net/url" @@ -845,6 +846,7 @@ func TestAddConfig_AllFieldsCovered(t *testing.T) { "ClientCertKeyPair": "non-config: parsed cert pair, not serialized", "Name": "non-config: profile name is not needed for debug purposes", "policy": "non-config: in-memory MDM policy snapshot, surfaced via Config.Policy() / GetConfigResponse.MDMManagedFields", + "probing": "non-config: marks a throwaway copy built to be diffed against; never set on a config anyone runs with", "DebugBundleUploadURL": "sensitive: MDM-provided upload URL may carry credentials or query tokens; kept out of the shared bundle", } @@ -969,3 +971,52 @@ func renderAddConfigSpecific(g *BundleGenerator) string { func newAnonymizerForTest() *anonymize.Anonymizer { return anonymize.NewAnonymizer(anonymize.DefaultAddresses()) } + +func TestRemoveStaleBundles(t *testing.T) { + dir := t.TempDir() + stale := filepath.Join(dir, "netbird.debug.111.zip") + fresh := filepath.Join(dir, "netbird.debug.222.zip") + other := filepath.Join(dir, "netbird.debug.333.txt") + owned := filepath.Join(dir, "netbird.debug.444.zip") + abandoned := filepath.Join(dir, "netbird.debug.555.zip") + for _, p := range []string{stale, fresh, other, owned, abandoned} { + require.NoError(t, os.WriteFile(p, []byte("x"), 0o600)) + } + exported, err := ExportBundle(owned) + require.NoError(t, err) + exportedAbandoned, err := ExportBundle(abandoned) + require.NoError(t, err) + old := time.Now().Add(-2 * time.Hour) + for _, p := range []string{stale, other, exported} { + require.NoError(t, os.Chtimes(p, old, old)) + } + ancient := time.Now().Add(-exportedBundleMaxAge - time.Hour) + require.NoError(t, os.Chtimes(exportedAbandoned, ancient, ancient)) + + RemoveStaleBundles(dir, time.Hour) + + assert.NoFileExists(t, stale, "bundle older than maxAge should be removed") + assert.FileExists(t, fresh, "bundle younger than maxAge must survive, it may still be uploading") + assert.FileExists(t, other, "files outside the bundle pattern must not be touched") + assert.NoFileExists(t, owned) + assert.FileExists(t, exported, "exported bundle is caller-owned and must survive maxAge") + assert.NoFileExists(t, exportedAbandoned, "exported bundle older than exportedBundleMaxAge is abandoned") +} + +func TestBundleIncludesNetworkMap(t *testing.T) { + for _, anonymize := range []bool{false, true} { + t.Run(fmt.Sprintf("anonymize=%t", anonymize), func(t *testing.T) { + g := NewBundleGenerator(GeneratorDependencies{ + SyncResponse: &mgmProto.SyncResponse{NetworkMap: &mgmProto.NetworkMap{Serial: 1}}, + }, BundleConfig{Anonymize: anonymize}) + + require.Contains(t, bundleEntries(t, g), "network_map.json") + }) + } +} + +func TestBundleOmitsNetworkMapWithoutSyncResponse(t *testing.T) { + g := NewBundleGenerator(GeneratorDependencies{}, BundleConfig{}) + + require.NotContains(t, bundleEntries(t, g), "network_map.json") +} diff --git a/client/internal/dns/host_windows.go b/client/internal/dns/host_windows.go index 948000a3d..6462d0c37 100644 --- a/client/internal/dns/host_windows.go +++ b/client/internal/dns/host_windows.go @@ -124,19 +124,9 @@ func newHostManager(wgInterface WGIface) (*registryConfigurator, error) { return nil, err } - var useGPO bool - k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE) - if err != nil { - log.Debugf("failed to open GPO DNS policy root: %v", err) - } else { - closer(k) - useGPO = true - log.Infof("detected GPO DNS policy configuration, using policy store") - } - configurator := ®istryConfigurator{ guid: guid, - gpo: useGPO, + gpo: useGPOPolicyStore(), } origNameservers, err := configurator.captureOriginalNameservers() @@ -576,14 +566,22 @@ func (r *registryConfigurator) setInterfaceRegistryKeyStringValue(key, value str return nil } +// deleteInterfaceRegistryKeyProperty removes a value from the interface key. +// A value that is already gone, or an interface key that is, is not an error: +// the caller asked for the value not to be there, and a cleanup that runs twice +// has to reach its later steps on the second run as well. func (r *registryConfigurator) deleteInterfaceRegistryKeyProperty(propertyKey string) error { regKey, err := r.getInterfaceRegistryKey() - if err != nil { + switch { + case errors.Is(err, registry.ErrNotExist), errors.Is(err, syscall.ERROR_PATH_NOT_FOUND): + log.Debugf("interface key of %s does not exist, nothing to delete %s from", r.guid, propertyKey) + return nil + case err != nil: return fmt.Errorf("get interface registry key: %w", err) } defer closer(regKey) - if err := regKey.DeleteValue(propertyKey); err != nil { + if err := regKey.DeleteValue(propertyKey); err != nil && !errors.Is(err, registry.ErrNotExist) { return fmt.Errorf("delete registry key %s: %w", propertyKey, err) } return nil @@ -612,7 +610,12 @@ func (r *registryConfigurator) restoreHostDNS() error { go r.flushDNSCache() - return nil + // Last, and only on the way out, once no rule of ours is left: during a + // session the store is where the rules of this run live, and emptying it + // mid-session would have the next rule recreate it anyway. Propagated so a + // failure keeps the shutdown state for the next run to retry, rather than + // leaving the store to hold up every rule change from here on. + return removeEmptyGPOPolicyStore() } // removeDNSMatchPolicies deletes every NRPT rule this client may have created, @@ -651,6 +654,73 @@ func (r *registryConfigurator) restoreUncleanShutdownDNS() error { return r.restoreHostDNS() } +// useGPOPolicyStore reports whether NRPT rules have to go into the group policy +// store, and clears an empty one out of the way first. +// +// The order is the point. A store left empty by an earlier run would otherwise +// decide this run too, sending its rules somewhere the resolver only reads when +// the policy engine next applies DNS client policy. Removing it before the +// choice is made leaves the local store authoritative for the whole session, +// including the first one after an upgrade. +func useGPOPolicyStore() bool { + if err := removeEmptyGPOPolicyStore(); err != nil { + // Nothing to retry against here: the worst case is the run going + // through the group policy store, which is where it would have gone + // before this check existed. + log.Warnf("%v", err) + } + + k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE) + if err != nil { + log.Debugf("failed to open GPO DNS policy root: %v", err) + return false + } + closer(k) + + log.Infof("detected GPO DNS policy configuration, using policy store") + return true +} + +// removeEmptyGPOPolicyStore deletes the group policy DnsPolicyConfig key once +// nothing is left in it. The key survives the deletion of the last rule it +// held, and the client treats its presence as "group policy configures the +// NRPT", so an empty one left behind keeps every later run writing rules there. +// Rules in that store reach the resolver only when the policy engine next +// applies DNS client policy, and a rule this client writes belongs to no GPO, +// so nothing schedules that application: both adding and removing a rule are +// held up by a minute or more, and for a removal that is a catch-all rule +// resolving every name over an interface that no longer exists. With the store +// absent the local one is authoritative and a change applies at once. +// +// A store that still holds rules, values or subkeys of somebody else's is left +// alone. +func removeEmptyGPOPolicyStore() error { + k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE) + switch { + case errors.Is(err, registry.ErrNotExist), errors.Is(err, syscall.ERROR_PATH_NOT_FOUND): + return nil + case err != nil: + return fmt.Errorf("open HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err) + } + + info, err := k.Stat() + closer(k) + if err != nil { + return fmt.Errorf("stat HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err) + } + + if info.SubKeyCount != 0 || info.ValueCount != 0 { + return nil + } + + if err := registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot); err != nil { + return fmt.Errorf("delete empty HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err) + } + + log.Infof("removed the empty GPO DNS policy store, leaving the local one authoritative") + return nil +} + // listNRPTRuleKeys returns the names of our NRPT rule keys under a policy store // root. An absent root holds nothing to clean up, which is the normal state of // the GPO store on a machine without DNS Client policy. diff --git a/client/internal/dns/host_windows_test.go b/client/internal/dns/host_windows_test.go index 7aef64590..353f6adbc 100644 --- a/client/internal/dns/host_windows_test.go +++ b/client/internal/dns/host_windows_test.go @@ -8,6 +8,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "golang.org/x/sys/windows/registry" + + "github.com/netbirdio/netbird/client/internal/winregistry" ) // TestNRPTEntriesCleanupOnConfigChange tests that old NRPT entries are properly cleaned up @@ -405,3 +407,130 @@ func TestNRPTDomainBatching(t *testing.T) { }) } } + +// TestRemoveEmptyGPOPolicyStore verifies that cleanup takes the GPO policy +// store itself with it once our rules are gone, since the store existing keeps +// the local one from being applied, and that a store with somebody else's rule +// in it is left alone. +func TestRemoveEmptyGPOPolicyStore(t *testing.T) { + if testing.Short() { + t.Skip("skipping registry integration test in short mode") + } + + t.Cleanup(func() { cleanupRegistryKeys(t) }) + cleanupRegistryKeys(t) + + testIP := netip.MustParseAddr("100.64.0.1") + cfg := ®istryConfigurator{gpo: true} + + // a store holding a rule of ours is kept, because the rule is still applied + require.NoError(t, cfg.addDNSMatchPolicy([]string{".example.com"}, testIP)) + exists, err := registryKeyExists(gpoDnsPolicyConfigMatchPath + "-0") + require.NoError(t, err) + require.True(t, exists, "Should write the rule to the GPO policy store") + + require.NoError(t, removeEmptyGPOPolicyStore()) + exists, err = registryKeyExists(GPODNSPolicyConfigRoot) + require.NoError(t, err) + assert.True(t, exists, "Should keep a policy store that still holds a rule") + + // once the rules are gone the store goes with them + require.NoError(t, cfg.removeDNSMatchPolicies()) + require.NoError(t, removeEmptyGPOPolicyStore()) + + exists, err = registryKeyExists(GPODNSPolicyConfigRoot) + require.NoError(t, err) + assert.False(t, exists, "Should remove the GPO policy store once it is empty") + + // A store is not ours to remove while somebody else has a rule in it. The + // rule is written volatile like our own: the rules above created the parent + // chain volatile, and Windows refuses a stable subkey under a volatile + // parent. + foreignRule := GPODNSPolicyConfigRoot + `\{2A3B4C5D-6E7F-4041-8283-84858687888A}` + foreignKey, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, foreignRule, registry.SET_VALUE) + require.NoError(t, err, "Should create a foreign GPO rule") + foreignKey.Close() + t.Cleanup(func() { + _ = registry.DeleteKey(registry.LOCAL_MACHINE, foreignRule) + _ = registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot) + }) + + require.NoError(t, cfg.removeDNSMatchPolicies()) + require.NoError(t, removeEmptyGPOPolicyStore()) + + exists, err = registryKeyExists(foreignRule) + require.NoError(t, err) + assert.True(t, exists, "Should not remove a foreign rule") + exists, err = registryKeyExists(GPODNSPolicyConfigRoot) + require.NoError(t, err) + assert.True(t, exists, "Should keep a policy store that still holds a foreign rule") +} + +// TestDeleteInterfaceRegistryKeyPropertyTwice verifies that removing a value +// that is already gone, or one on an interface key that is, reports success. +// Teardown runs again after a failed cleanup, and the steps that follow this +// one have to be reached on that second run. +func TestDeleteInterfaceRegistryKeyPropertyTwice(t *testing.T) { + if testing.Short() { + t.Skip("skipping registry integration test in short mode") + } + + testGUID := "{12345678-1234-1234-1234-123456789ABC}" + interfacePath := InterfaceConfigPath + `\` + testGUID + testKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, interfacePath, registry.SET_VALUE) + require.NoError(t, err, "Should create test interface registry key") + testKey.Close() + t.Cleanup(func() { + _ = registry.DeleteKey(registry.LOCAL_MACHINE, interfacePath) + }) + + cfg := ®istryConfigurator{guid: testGUID} + + require.NoError(t, cfg.setInterfaceRegistryKeyStringValue(interfaceConfigSearchListKey, "example.com")) + require.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey)) + assert.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey), + "Should report success for a value that is already gone") + + // and with the interface key itself gone, as it is once the adapter is + require.NoError(t, registry.DeleteKey(registry.LOCAL_MACHINE, interfacePath)) + assert.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey), + "Should report success when the interface key does not exist") +} + +// TestUseGPOPolicyStoreClearsEmptyStore verifies that the store is cleared +// before it is consulted, so an empty one left by an earlier run does not send +// this run's rules to the group policy store. A store somebody else has a rule +// in still decides where the rules go. +func TestUseGPOPolicyStoreClearsEmptyStore(t *testing.T) { + if testing.Short() { + t.Skip("skipping registry integration test in short mode") + } + + t.Cleanup(func() { cleanupRegistryKeys(t) }) + cleanupRegistryKeys(t) + + // the leftover an earlier run used to keep, which the client read as + // "group policy configures the NRPT" for every run after it + emptyStore, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.SET_VALUE) + require.NoError(t, err, "Should create the GPO policy store") + emptyStore.Close() + + assert.False(t, useGPOPolicyStore(), "An empty store should not decide where the rules go") + exists, err := registryKeyExists(GPODNSPolicyConfigRoot) + require.NoError(t, err) + assert.False(t, exists, "Should clear the empty store before consulting it") + + foreignRule := GPODNSPolicyConfigRoot + `\{2A3B4C5D-6E7F-4041-8283-84858687888A}` + foreignKey, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, foreignRule, registry.SET_VALUE) + require.NoError(t, err, "Should create a foreign GPO rule") + foreignKey.Close() + t.Cleanup(func() { + _ = registry.DeleteKey(registry.LOCAL_MACHINE, foreignRule) + _ = registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot) + }) + + assert.True(t, useGPOPolicyStore(), "A store holding a rule should decide where the rules go") + exists, err = registryKeyExists(GPODNSPolicyConfigRoot) + require.NoError(t, err) + assert.True(t, exists, "Should keep a store that holds a rule") +} diff --git a/client/internal/dns/server_privileged_test.go b/client/internal/dns/server_privileged_test.go index a17044cf5..b135c3661 100644 --- a/client/internal/dns/server_privileged_test.go +++ b/client/internal/dns/server_privileged_test.go @@ -9,9 +9,9 @@ import ( "os" "testing" - "go.uber.org/mock/gomock" "github.com/miekg/dns" "github.com/stretchr/testify/assert" + "go.uber.org/mock/gomock" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" "github.com/netbirdio/netbird/client/iface" @@ -24,6 +24,10 @@ import ( nbdns "github.com/netbirdio/netbird/dns" ) +// testIFaceBlackList mirrors the overlay prefixes profilemanager.DefaultInterfaceBlacklist +// carries. Declared here rather than imported because profilemanager imports this package. +var testIFaceBlackList = []string{"wt", "utun", "tun0"} + func TestUpdateDNSServer(t *testing.T) { nameServers := []nbdns.NameServer{ @@ -243,10 +247,7 @@ func TestUpdateDNSServer(t *testing.T) { for n, testCase := range testCases { t.Run(testCase.name, func(t *testing.T) { privKey, _ := wgtypes.GenerateKey() - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil) opts := iface.WGIFaceOpts{ IFaceName: fmt.Sprintf("utun230%d", n), @@ -348,11 +349,7 @@ func TestDNSFakeResolverHandleUpdates(t *testing.T) { defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) t.Setenv("NB_WG_KERNEL_DISABLED", "true") - newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"}) - if err != nil { - t.Errorf("create stdnet: %v", err) - return - } + newNet := stdnet.NewNet(context.Background(), []string{"utun2301"}, nil) privKey, _ := wgtypes.GeneratePrivateKey() opts := iface.WGIFaceOpts{ diff --git a/client/internal/dns/server_test.go b/client/internal/dns/server_test.go index 0144a4a8b..777594272 100644 --- a/client/internal/dns/server_test.go +++ b/client/internal/dns/server_test.go @@ -394,11 +394,7 @@ func createWgInterfaceWithBind(t *testing.T) (*iface.WGIface, error) { defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) t.Setenv("NB_WG_KERNEL_DISABLED", "true") - newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"}) - if err != nil { - t.Fatalf("create stdnet: %v", err) - return nil, err - } + newNet := stdnet.NewNet(context.Background(), []string{"utun2301"}, nil) privKey, _ := wgtypes.GeneratePrivateKey() diff --git a/client/internal/dns/service_listener.go b/client/internal/dns/service_listener.go index 3dc29c4dc..d65a727b1 100644 --- a/client/internal/dns/service_listener.go +++ b/client/internal/dns/service_listener.go @@ -6,6 +6,7 @@ import ( "net" "net/netip" "runtime" + "slices" "strconv" "sync" "time" @@ -17,17 +18,20 @@ import ( nberrors "github.com/netbirdio/netbird/client/errors" firewall "github.com/netbirdio/netbird/client/firewall/manager" - "github.com/netbirdio/netbird/client/internal/ebpf" - ebpfMgr "github.com/netbirdio/netbird/client/internal/ebpf/manager" ) const ( customPort = 5053 + // randomPortAttempts bounds the search for a port free on both protocols. + randomPortAttempts = 5 ) var ( defaultIP = netip.MustParseAddr("127.0.0.1") customIP = netip.MustParseAddr("127.0.0.153") + + // dnatProtocols are the protocols the port 53 redirect covers. + dnatProtocols = []firewall.Protocol{firewall.ProtocolUDP, firewall.ProtocolTCP} ) type serviceViaListener struct { @@ -40,9 +44,20 @@ type serviceViaListener struct { listenPort uint16 listenerIsRunning bool listenerFlagLock sync.Mutex - ebpfService ebpfMgr.Manager firewall Firewall - tcpDNATConfigured bool + // dnatRules holds the port 53 redirects that are installed and not yet + // removed, so a removal that fails can be retried. + dnatRules []dnatRule +} + +// dnatRule is a port 53 redirect as it was installed. The target is kept with +// the rule because the listener can come back on a different address or port, +// and a retried removal has to name the address and port the rule was added +// with, not the ones in use now. +type dnatRule struct { + protocol firewall.Protocol + ip netip.Addr + port uint16 } func newServiceViaListener(wgIface WGIface, customAddr *netip.AddrPort, fw Firewall) *serviceViaListener { @@ -112,34 +127,93 @@ func (s *serviceViaListener) Listen() error { } }() - // When eBPF redirects UDP port 53 to our listen port, TCP still needs - // a DNAT rule because eBPF only handles UDP. - if s.ebpfService != nil && s.firewall != nil && s.listenPort != DefaultPort { - if err := s.firewall.AddOutputDNAT(s.listenIP, firewall.ProtocolTCP, DefaultPort, s.listenPort); err != nil { - log.Warnf("failed to add DNS TCP DNAT rule, TCP DNS on port 53 will not work: %v", err) - } else { - s.tcpDNATConfigured = true - log.Infof("added DNS TCP DNAT rule: %s:%d -> %s:%d", s.listenIP, DefaultPort, s.listenIP, s.listenPort) - } + if s.listenPort != DefaultPort { + s.setupDNAT() } return nil } +// setupDNAT redirects port 53 to the port the DNS server actually listens on. +// Both protocols must be redirected or none: RuntimePort reports port 53 only +// while the full redirect is in place, so a half-configured redirect would +// advertise a resolver that answers over one protocol. +func (s *serviceViaListener) setupDNAT() { + if s.firewall == nil { + log.Errorf("no firewall manager available to redirect DNS port %d to %d, "+ + "clients pointed at %s will not reach the resolver", DefaultPort, s.listenPort, s.listenIP) + return + } + + // Clear whatever an earlier removal left behind first. Those rules can point + // at an address or port this listener no longer uses, and they are matched + // before anything added now, so adding a redirect on top of one would keep + // sending port 53 traffic to the previous listener while reporting the + // redirect as complete. The rules stay recorded for a later attempt. + if err := s.removeDNAT(); err != nil { + log.Errorf("failed to remove stale DNS DNAT rules, leaving port %d redirected to the previous listener: %v", + DefaultPort, err) + return + } + + for _, proto := range dnatProtocols { + if err := s.firewall.AddOutputDNAT(s.listenIP, proto, DefaultPort, s.listenPort); err != nil { + log.Errorf("failed to add DNS %s DNAT rule, DNS on port %d will not work: %v", + proto, DefaultPort, err) + if err := s.removeDNAT(); err != nil { + log.Warnf("failed to roll back DNS DNAT rules, retrying on stop: %v", err) + } + return + } + s.dnatRules = append(s.dnatRules, dnatRule{protocol: proto, ip: s.listenIP, port: s.listenPort}) + } + + log.Infof("added DNS DNAT rules: %s:%d -> %s:%d (UDP + TCP)", s.listenIP, DefaultPort, s.listenIP, s.listenPort) +} + +// removeDNAT removes every installed port 53 redirect. A rule whose removal +// fails stays recorded so a later setup or Stop retries it, rather than leaving +// port 53 pointing at a resolver that is no longer listening. +func (s *serviceViaListener) removeDNAT() error { + if s.firewall == nil { + return nil + } + + var merr *multierror.Error + var remaining []dnatRule + for _, rule := range s.dnatRules { + if err := s.firewall.RemoveOutputDNAT(rule.ip, rule.protocol, DefaultPort, rule.port); err != nil { + merr = multierror.Append(merr, fmt.Errorf("remove DNS %s DNAT rule for %s:%d: %w", + rule.protocol, rule.ip, rule.port, err)) + remaining = append(remaining, rule) + } + } + s.dnatRules = remaining + + return nberrors.FormatErrorOrNil(merr) +} + func (s *serviceViaListener) Stop() error { s.listenerFlagLock.Lock() defer s.listenerFlagLock.Unlock() + var merr *multierror.Error + + // Redirects are removed even when the listener is already stopped, so that + // a removal which failed earlier is retried instead of leaving port 53 + // pointing at a resolver that no longer listens. + if err := s.removeDNAT(); err != nil { + merr = multierror.Append(merr, err) + } + if !s.listenerIsRunning { - return nil + return nberrors.FormatErrorOrNil(merr) } s.listenerIsRunning = false ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() - var merr *multierror.Error - if err := s.server.ShutdownContext(ctx); err != nil { merr = multierror.Append(merr, fmt.Errorf("stop DNS UDP server: %w", err)) } @@ -148,19 +222,6 @@ func (s *serviceViaListener) Stop() error { merr = multierror.Append(merr, fmt.Errorf("stop DNS TCP server: %w", err)) } - if s.tcpDNATConfigured && s.firewall != nil { - if err := s.firewall.RemoveOutputDNAT(s.listenIP, firewall.ProtocolTCP, DefaultPort, s.listenPort); err != nil { - merr = multierror.Append(merr, fmt.Errorf("remove DNS TCP DNAT rule: %w", err)) - } - s.tcpDNATConfigured = false - } - - if s.ebpfService != nil { - if err := s.ebpfService.FreeDNSFwd(); err != nil { - merr = multierror.Append(merr, fmt.Errorf("stop traffic forwarder: %w", err)) - } - } - return nberrors.FormatErrorOrNil(merr) } @@ -177,11 +238,23 @@ func (s *serviceViaListener) RuntimePort() int { s.listenerFlagLock.Lock() defer s.listenerFlagLock.Unlock() - if s.ebpfService != nil { + if s.redirectInstalled() { return DefaultPort - } else { - return int(s.listenPort) } + return int(s.listenPort) +} + +// redirectInstalled reports whether every protocol is redirected from port 53 +// to the address and port the listener currently serves. Rules left over from +// an earlier listener do not count. +func (s *serviceViaListener) redirectInstalled() bool { + for _, proto := range dnatProtocols { + current := dnatRule{protocol: proto, ip: s.listenIP, port: s.listenPort} + if !slices.Contains(s.dnatRules, current) { + return false + } + } + return true } func (s *serviceViaListener) RuntimeIP() netip.Addr { @@ -190,30 +263,29 @@ func (s *serviceViaListener) RuntimeIP() netip.Addr { // evalListenAddress figures out the listen address for the DNS server. // IPv4-only: all peers have a v4 overlay address, and DNS config points to v4. -// First checks port 53 on WG interface or lo, then tries eBPF on a random port, -// then falls back to port 5053. +// Prefers port 53 on the overlay interface or lo, so no redirect is needed at +// all; when it is taken it falls back to port 5053 and then to a random free +// port, both of which need the port 53 redirect set up by setupDNAT. func (s *serviceViaListener) evalListenAddress() (netip.Addr, uint16, error) { if s.customAddr != nil { return s.customAddr.Addr(), s.customAddr.Port(), nil } - ip, ok := s.testFreePort(DefaultPort) - if ok { + if ip, ok := s.testFreePort(DefaultPort); ok { return ip, DefaultPort, nil } - ebpfSrv, port, ok := s.tryToUseeBPF() - if ok { - s.ebpfService = ebpfSrv - return s.wgInterface.Address().IP, port, nil - } - - ip, ok = s.testFreePort(customPort) - if ok { + if ip, ok := s.testFreePort(customPort); ok { return ip, customPort, nil } - return netip.Addr{}, 0, fmt.Errorf("failed to find a free port for DNS server") + ip := s.wgInterface.Address().IP + port, err := s.randomFreePort(ip) + if err != nil { + return netip.Addr{}, 0, fmt.Errorf("find a free port for DNS server: %w", err) + } + + return ip, port, nil } func (s *serviceViaListener) testFreePort(port int) (netip.Addr, bool) { @@ -260,48 +332,25 @@ func (s *serviceViaListener) tryToBind(ip netip.Addr, port int) bool { return true } -// tryToUseeBPF decides whether to apply eBPF program to capture DNS traffic on port 53. -// This is needed because on some operating systems if we start a DNS server not on a default port 53, -// the domain name resolution won't work. So, in case we are running on Linux and picked a free -// port we should fall back to the eBPF solution that will capture traffic on port 53 and forward -// it to a local DNS server running on the chosen port. -func (s *serviceViaListener) tryToUseeBPF() (ebpfMgr.Manager, uint16, bool) { - if runtime.GOOS != "linux" { - return nil, 0, false +// randomFreePort returns a port that is free on ip for both UDP and TCP, since +// the DNS server binds both. The probe listeners are closed again, so the port +// is only likely, not guaranteed, to still be free when the server binds it. +func (s *serviceViaListener) randomFreePort(ip netip.Addr) (uint16, error) { + for range randomPortAttempts { + probeListener, err := net.ListenUDP("udp4", &net.UDPAddr{}) + if err != nil { + return 0, fmt.Errorf("bind random port: %w", err) + } + + port := uint16(probeListener.LocalAddr().(*net.UDPAddr).Port) + if err := probeListener.Close(); err != nil { + return 0, fmt.Errorf("free up probed port: %w", err) + } + + if s.tryToBind(ip, int(port)) { + return port, nil + } } - port, err := s.generateFreePort() //nolint:staticcheck,unused - if err != nil { - log.Warnf("failed to generate a free port for eBPF DNS forwarder server: %s", err) - return nil, 0, false - } - - ebpfSrv := ebpf.GetEbpfManagerInstance() - err = ebpfSrv.LoadDNSFwd(s.wgInterface.Address().IP, int(port)) - if err != nil { - log.Warnf("failed to load DNS forwarder eBPF program, error: %s", err) - return nil, 0, false - } - - return ebpfSrv, port, true -} - -func (s *serviceViaListener) generateFreePort() (uint16, error) { - ok := s.tryToBind(s.wgInterface.Address().IP, customPort) - if ok { - return customPort, nil - } - - probeListener, err := net.ListenUDP("udp4", &net.UDPAddr{}) - if err != nil { - log.Debugf("failed to bind random port for DNS: %s", err) - return 0, err - } - - port := uint16(probeListener.LocalAddr().(*net.UDPAddr).Port) - if err = probeListener.Close(); err != nil { - log.Debugf("failed to free up DNS port: %s", err) - return 0, err - } - return port, nil + return 0, fmt.Errorf("no port free for UDP and TCP on %s after %d attempts", ip, randomPortAttempts) } diff --git a/client/internal/dns/service_listener_test.go b/client/internal/dns/service_listener_test.go index 90ef71d19..b158a79fd 100644 --- a/client/internal/dns/service_listener_test.go +++ b/client/internal/dns/service_listener_test.go @@ -1,6 +1,7 @@ package dns import ( + "errors" "fmt" "net" "net/netip" @@ -10,6 +11,8 @@ import ( "github.com/miekg/dns" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + firewall "github.com/netbirdio/netbird/client/firewall/manager" ) func TestServiceViaListener_TCPAndUDP(t *testing.T) { @@ -84,3 +87,133 @@ func TestServiceViaListener_TCPAndUDP(t *testing.T) { require.NotEmpty(t, tcpResp.Answer) assert.Contains(t, tcpResp.Answer[0].String(), "192.0.2.1", "TCP response should contain expected IP") } + +type dnatCall struct { + rule dnatRule + added bool +} + +// fakeFirewall records DNAT calls and fails the ones named in addErrs/removeErrs. +type fakeFirewall struct { + calls []dnatCall + addErrs map[firewall.Protocol]error + removeErrs map[firewall.Protocol]error +} + +func (f *fakeFirewall) AddOutputDNAT(ip netip.Addr, protocol firewall.Protocol, _, translatedPort uint16) error { + if err := f.addErrs[protocol]; err != nil { + return err + } + f.calls = append(f.calls, dnatCall{rule: dnatRule{protocol: protocol, ip: ip, port: translatedPort}, added: true}) + return nil +} + +func (f *fakeFirewall) RemoveOutputDNAT(ip netip.Addr, protocol firewall.Protocol, _, translatedPort uint16) error { + if err := f.removeErrs[protocol]; err != nil { + return err + } + f.calls = append(f.calls, dnatCall{rule: dnatRule{protocol: protocol, ip: ip, port: translatedPort}}) + return nil +} + +func newDNATTestService(fw Firewall) *serviceViaListener { + return &serviceViaListener{ + listenIP: netip.MustParseAddr("100.64.0.1"), + listenPort: customPort, + firewall: fw, + } +} + +func TestSetupDNAT_BothProtocols(t *testing.T) { + svc := newDNATTestService(&fakeFirewall{}) + + svc.setupDNAT() + + assert.Len(t, svc.dnatRules, len(dnatProtocols)) + assert.Equal(t, DefaultPort, svc.RuntimePort(), "port 53 is advertised once both redirects are installed") +} + +func TestSetupDNAT_RollsBackPartialRedirect(t *testing.T) { + fw := &fakeFirewall{addErrs: map[firewall.Protocol]error{firewall.ProtocolTCP: errors.New("nftables busy")}} + svc := newDNATTestService(fw) + + svc.setupDNAT() + + assert.Empty(t, svc.dnatRules, "the UDP redirect installed before the failure must be rolled back") + assert.Equal(t, int(svc.listenPort), svc.RuntimePort(), "an incomplete redirect must not advertise port 53") + udp := dnatRule{protocol: firewall.ProtocolUDP, ip: svc.listenIP, port: svc.listenPort} + assert.Contains(t, fw.calls, dnatCall{rule: udp}, "UDP removal should have been attempted") +} + +// A rollback that fails must keep the rule recorded, so port 53 is not left +// redirected to a resolver that no longer listens. +func TestStop_RetriesFailedDNATRemoval(t *testing.T) { + fw := &fakeFirewall{ + addErrs: map[firewall.Protocol]error{firewall.ProtocolTCP: errors.New("nftables busy")}, + removeErrs: map[firewall.Protocol]error{firewall.ProtocolUDP: errors.New("nftables busy")}, + } + svc := newDNATTestService(fw) + + svc.setupDNAT() + udp := dnatRule{protocol: firewall.ProtocolUDP, ip: svc.listenIP, port: svc.listenPort} + require.Equal(t, []dnatRule{udp}, svc.dnatRules, "a failed rollback keeps the rule for a later retry") + + require.Error(t, svc.Stop(), "the failing removal should be reported") + require.Equal(t, []dnatRule{udp}, svc.dnatRules) + + delete(fw.removeErrs, firewall.ProtocolUDP) + require.NoError(t, svc.Stop(), "a later stop retries the removal") + assert.Empty(t, svc.dnatRules) +} + +// A stale rule that cannot be removed is matched before anything added now, so +// no new redirect may be installed on top of it and port 53 must not be +// advertised as reaching this listener. +func TestSetupDNAT_AbortsWhileStaleRuleRemains(t *testing.T) { + fw := &fakeFirewall{removeErrs: map[firewall.Protocol]error{firewall.ProtocolUDP: errors.New("nftables busy")}} + svc := newDNATTestService(fw) + stalePort := svc.listenPort + + svc.setupDNAT() + require.Error(t, svc.Stop()) + staleUDP := dnatRule{protocol: firewall.ProtocolUDP, ip: svc.listenIP, port: stalePort} + require.Equal(t, []dnatRule{staleUDP}, svc.dnatRules) + + svc.listenPort = stalePort + 1 + fw.calls = nil + + svc.setupDNAT() + + assert.Equal(t, []dnatRule{staleUDP}, svc.dnatRules, "the stale rule stays recorded for a later attempt") + for _, call := range fw.calls { + assert.False(t, call.added, "no redirect may be installed while a stale one is still in place") + } + assert.Equal(t, int(svc.listenPort), svc.RuntimePort(), "port 53 must not be advertised") +} + +// A rule left behind by a failed removal must be removed with the address and +// port it was installed with, even when the listener has since moved to another +// port, and it must not count towards the redirect the new listener advertises. +func TestSetupDNAT_ClearsStaleRuleAfterPortChange(t *testing.T) { + fw := &fakeFirewall{removeErrs: map[firewall.Protocol]error{firewall.ProtocolUDP: errors.New("nftables busy")}} + svc := newDNATTestService(fw) + stalePort := svc.listenPort + + svc.setupDNAT() + require.Error(t, svc.Stop()) + staleUDP := dnatRule{protocol: firewall.ProtocolUDP, ip: svc.listenIP, port: stalePort} + require.Equal(t, []dnatRule{staleUDP}, svc.dnatRules) + + delete(fw.removeErrs, firewall.ProtocolUDP) + svc.listenPort = stalePort + 1 + fw.calls = nil + + svc.setupDNAT() + + assert.Contains(t, fw.calls, dnatCall{rule: staleUDP}, "the stale rule must be removed with its original port") + assert.Len(t, svc.dnatRules, len(dnatProtocols)) + assert.Equal(t, DefaultPort, svc.RuntimePort(), "the new listener is fully redirected") + for _, rule := range svc.dnatRules { + assert.Equal(t, svc.listenPort, rule.port, "only rules for the current listener remain") + } +} diff --git a/client/internal/ebpf/ebpf/bpf_bpfeb.go b/client/internal/ebpf/ebpf/bpf_bpfeb.go deleted file mode 100644 index 04b19883b..000000000 --- a/client/internal/ebpf/ebpf/bpf_bpfeb.go +++ /dev/null @@ -1,128 +0,0 @@ -// Code generated by bpf2go; DO NOT EDIT. -//go:build arm64be || armbe || mips || mips64 || mips64p32 || ppc64 || s390 || s390x || sparc || sparc64 - -package ebpf - -import ( - "bytes" - _ "embed" - "fmt" - "io" - - "github.com/cilium/ebpf" -) - -// loadBpf returns the embedded CollectionSpec for bpf. -func loadBpf() (*ebpf.CollectionSpec, error) { - reader := bytes.NewReader(_BpfBytes) - spec, err := ebpf.LoadCollectionSpecFromReader(reader) - if err != nil { - return nil, fmt.Errorf("can't load bpf: %w", err) - } - - return spec, err -} - -// loadBpfObjects loads bpf and converts it into a struct. -// -// The following types are suitable as obj argument: -// -// *bpfObjects -// *bpfPrograms -// *bpfMaps -// -// See ebpf.CollectionSpec.LoadAndAssign documentation for details. -func loadBpfObjects(obj interface{}, opts *ebpf.CollectionOptions) error { - spec, err := loadBpf() - if err != nil { - return err - } - - return spec.LoadAndAssign(obj, opts) -} - -// bpfSpecs contains maps and programs before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfSpecs struct { - bpfProgramSpecs - bpfMapSpecs -} - -// bpfSpecs contains programs before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfProgramSpecs struct { - NbXdpProg *ebpf.ProgramSpec `ebpf:"nb_xdp_prog"` -} - -// bpfMapSpecs contains maps before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfMapSpecs struct { - NbFeatures *ebpf.MapSpec `ebpf:"nb_features"` - NbMapDnsIp *ebpf.MapSpec `ebpf:"nb_map_dns_ip"` - NbMapDnsPort *ebpf.MapSpec `ebpf:"nb_map_dns_port"` - NbWgProxySettingsMap *ebpf.MapSpec `ebpf:"nb_wg_proxy_settings_map"` -} - -// bpfObjects contains all objects after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfObjects struct { - bpfPrograms - bpfMaps -} - -func (o *bpfObjects) Close() error { - return _BpfClose( - &o.bpfPrograms, - &o.bpfMaps, - ) -} - -// bpfMaps contains all maps after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfMaps struct { - NbFeatures *ebpf.Map `ebpf:"nb_features"` - NbMapDnsIp *ebpf.Map `ebpf:"nb_map_dns_ip"` - NbMapDnsPort *ebpf.Map `ebpf:"nb_map_dns_port"` - NbWgProxySettingsMap *ebpf.Map `ebpf:"nb_wg_proxy_settings_map"` -} - -func (m *bpfMaps) Close() error { - return _BpfClose( - m.NbFeatures, - m.NbMapDnsIp, - m.NbMapDnsPort, - m.NbWgProxySettingsMap, - ) -} - -// bpfPrograms contains all programs after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfPrograms struct { - NbXdpProg *ebpf.Program `ebpf:"nb_xdp_prog"` -} - -func (p *bpfPrograms) Close() error { - return _BpfClose( - p.NbXdpProg, - ) -} - -func _BpfClose(closers ...io.Closer) error { - for _, closer := range closers { - if err := closer.Close(); err != nil { - return err - } - } - return nil -} - -// Do not access this directly. -// -//go:embed bpf_bpfeb.o -var _BpfBytes []byte diff --git a/client/internal/ebpf/ebpf/bpf_bpfeb.o b/client/internal/ebpf/ebpf/bpf_bpfeb.o deleted file mode 100644 index 7433ad740..000000000 Binary files a/client/internal/ebpf/ebpf/bpf_bpfeb.o and /dev/null differ diff --git a/client/internal/ebpf/ebpf/bpf_bpfel.go b/client/internal/ebpf/ebpf/bpf_bpfel.go deleted file mode 100644 index 03b494aa2..000000000 --- a/client/internal/ebpf/ebpf/bpf_bpfel.go +++ /dev/null @@ -1,128 +0,0 @@ -// Code generated by bpf2go; DO NOT EDIT. -//go:build 386 || amd64 || amd64p32 || arm || arm64 || loong64 || mips64le || mips64p32le || mipsle || ppc64le || riscv64 - -package ebpf - -import ( - "bytes" - _ "embed" - "fmt" - "io" - - "github.com/cilium/ebpf" -) - -// loadBpf returns the embedded CollectionSpec for bpf. -func loadBpf() (*ebpf.CollectionSpec, error) { - reader := bytes.NewReader(_BpfBytes) - spec, err := ebpf.LoadCollectionSpecFromReader(reader) - if err != nil { - return nil, fmt.Errorf("can't load bpf: %w", err) - } - - return spec, err -} - -// loadBpfObjects loads bpf and converts it into a struct. -// -// The following types are suitable as obj argument: -// -// *bpfObjects -// *bpfPrograms -// *bpfMaps -// -// See ebpf.CollectionSpec.LoadAndAssign documentation for details. -func loadBpfObjects(obj interface{}, opts *ebpf.CollectionOptions) error { - spec, err := loadBpf() - if err != nil { - return err - } - - return spec.LoadAndAssign(obj, opts) -} - -// bpfSpecs contains maps and programs before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfSpecs struct { - bpfProgramSpecs - bpfMapSpecs -} - -// bpfSpecs contains programs before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfProgramSpecs struct { - NbXdpProg *ebpf.ProgramSpec `ebpf:"nb_xdp_prog"` -} - -// bpfMapSpecs contains maps before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfMapSpecs struct { - NbFeatures *ebpf.MapSpec `ebpf:"nb_features"` - NbMapDnsIp *ebpf.MapSpec `ebpf:"nb_map_dns_ip"` - NbMapDnsPort *ebpf.MapSpec `ebpf:"nb_map_dns_port"` - NbWgProxySettingsMap *ebpf.MapSpec `ebpf:"nb_wg_proxy_settings_map"` -} - -// bpfObjects contains all objects after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfObjects struct { - bpfPrograms - bpfMaps -} - -func (o *bpfObjects) Close() error { - return _BpfClose( - &o.bpfPrograms, - &o.bpfMaps, - ) -} - -// bpfMaps contains all maps after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfMaps struct { - NbFeatures *ebpf.Map `ebpf:"nb_features"` - NbMapDnsIp *ebpf.Map `ebpf:"nb_map_dns_ip"` - NbMapDnsPort *ebpf.Map `ebpf:"nb_map_dns_port"` - NbWgProxySettingsMap *ebpf.Map `ebpf:"nb_wg_proxy_settings_map"` -} - -func (m *bpfMaps) Close() error { - return _BpfClose( - m.NbFeatures, - m.NbMapDnsIp, - m.NbMapDnsPort, - m.NbWgProxySettingsMap, - ) -} - -// bpfPrograms contains all programs after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfPrograms struct { - NbXdpProg *ebpf.Program `ebpf:"nb_xdp_prog"` -} - -func (p *bpfPrograms) Close() error { - return _BpfClose( - p.NbXdpProg, - ) -} - -func _BpfClose(closers ...io.Closer) error { - for _, closer := range closers { - if err := closer.Close(); err != nil { - return err - } - } - return nil -} - -// Do not access this directly. -// -//go:embed bpf_bpfel.o -var _BpfBytes []byte diff --git a/client/internal/ebpf/ebpf/bpf_bpfel.o b/client/internal/ebpf/ebpf/bpf_bpfel.o deleted file mode 100644 index 779f43a00..000000000 Binary files a/client/internal/ebpf/ebpf/bpf_bpfel.o and /dev/null differ diff --git a/client/internal/ebpf/ebpf/dns_fwd_linux.go b/client/internal/ebpf/ebpf/dns_fwd_linux.go deleted file mode 100644 index 1e7774573..000000000 --- a/client/internal/ebpf/ebpf/dns_fwd_linux.go +++ /dev/null @@ -1,52 +0,0 @@ -package ebpf - -import ( - "encoding/binary" - "fmt" - "net/netip" - - log "github.com/sirupsen/logrus" -) - -const ( - mapKeyDNSIP uint32 = 0 - mapKeyDNSPort uint32 = 1 -) - -func (tf *GeneralManager) LoadDNSFwd(ip netip.Addr, dnsPort int) error { - log.Debugf("load eBPF DNS forwarder, watching addr: %s:53, redirect to port: %d", ip, dnsPort) - tf.lock.Lock() - defer tf.lock.Unlock() - - err := tf.loadXdp() - if err != nil { - return err - } - - if !ip.Is4() { - return fmt.Errorf("eBPF DNS forwarder only supports IPv4, got %s", ip) - } - ip4 := ip.As4() - err = tf.bpfObjs.NbMapDnsIp.Put(mapKeyDNSIP, binary.BigEndian.Uint32(ip4[:])) - if err != nil { - return err - } - - err = tf.bpfObjs.NbMapDnsPort.Put(mapKeyDNSPort, uint16(dnsPort)) - if err != nil { - return err - } - - tf.setFeatureFlag(featureFlagDnsForwarder) - err = tf.bpfObjs.NbFeatures.Put(mapKeyFeatures, tf.featureFlags) - if err != nil { - return err - } - return nil -} - -func (tf *GeneralManager) FreeDNSFwd() error { - log.Debugf("free ebpf DNS forwarder") - return tf.unsetFeatureFlag(featureFlagDnsForwarder) -} - diff --git a/client/internal/ebpf/ebpf/manager_linux.go b/client/internal/ebpf/ebpf/manager_linux.go deleted file mode 100644 index 7520a6387..000000000 --- a/client/internal/ebpf/ebpf/manager_linux.go +++ /dev/null @@ -1,116 +0,0 @@ -package ebpf - -import ( - _ "embed" - "net" - "sync" - - "github.com/cilium/ebpf/link" - "github.com/cilium/ebpf/rlimit" - log "github.com/sirupsen/logrus" - - "github.com/netbirdio/netbird/client/internal/ebpf/manager" -) - -const ( - mapKeyFeatures uint32 = 0 - - featureFlagWGProxy = 0b00000001 - featureFlagDnsForwarder = 0b00000010 -) - -var ( - singleton manager.Manager - singletonLock = &sync.Mutex{} -) - -// required packages libbpf-dev, libc6-dev-i386-amd64-cross - -// GeneralManager is used to load multiple eBPF programs with a custom check (if then) done in prog.c -// The manager simply adds a feature (byte) of each program to a map that is shared between the userspace and kernel. -// When packet arrives, the C code checks for each feature (if it is set) and executes each enabled program (e.g., dns_fwd.c and wg_proxy.c). -// -//go:generate go run github.com/cilium/ebpf/cmd/bpf2go -cc clang-14 bpf src/prog.c -- -I /usr/x86_64-linux-gnu/include -type GeneralManager struct { - lock sync.Mutex - link link.Link - featureFlags uint16 - bpfObjs bpfObjects -} - -// GetEbpfManagerInstance return a static eBpf Manager instance -func GetEbpfManagerInstance() manager.Manager { - singletonLock.Lock() - defer singletonLock.Unlock() - if singleton != nil { - return singleton - } - singleton = &GeneralManager{} - return singleton -} - -func (tf *GeneralManager) setFeatureFlag(feature uint16) { - tf.featureFlags |= feature -} - -func (tf *GeneralManager) loadXdp() error { - if tf.link != nil { - return nil - } - // it required for Docker - err := rlimit.RemoveMemlock() - if err != nil { - return err - } - - iFace, err := net.InterfaceByName("lo") - if err != nil { - return err - } - - // load pre-compiled programs into the kernel. - err = loadBpfObjects(&tf.bpfObjs, nil) - if err != nil { - return err - } - - tf.link, err = link.AttachXDP(link.XDPOptions{ - Program: tf.bpfObjs.NbXdpProg, - Interface: iFace.Index, - }) - - if err != nil { - _ = tf.bpfObjs.Close() - tf.link = nil - return err - } - return nil -} - -func (tf *GeneralManager) unsetFeatureFlag(feature uint16) error { - tf.lock.Lock() - defer tf.lock.Unlock() - tf.featureFlags &^= feature - - if tf.link == nil { - return nil - } - - if tf.featureFlags == 0 { - return tf.close() - } - - return tf.bpfObjs.NbFeatures.Put(mapKeyFeatures, tf.featureFlags) -} - -func (tf *GeneralManager) close() error { - log.Debugf("detach ebpf program ") - err := tf.bpfObjs.Close() - if err != nil { - log.Warnf("failed to close eBpf objects: %s", err) - } - - err = tf.link.Close() - tf.link = nil - return err -} diff --git a/client/internal/ebpf/ebpf/manager_linux_test.go b/client/internal/ebpf/ebpf/manager_linux_test.go deleted file mode 100644 index 5664a4565..000000000 --- a/client/internal/ebpf/ebpf/manager_linux_test.go +++ /dev/null @@ -1,40 +0,0 @@ -package ebpf - -import ( - "testing" -) - -func TestManager_setFeatureFlag(t *testing.T) { - mgr := GeneralManager{} - mgr.setFeatureFlag(featureFlagWGProxy) - if mgr.featureFlags != 1 { - t.Errorf("invalid feature state") - } - - mgr.setFeatureFlag(featureFlagDnsForwarder) - if mgr.featureFlags != 3 { - t.Errorf("invalid feature state") - } -} - -func TestManager_unsetFeatureFlag(t *testing.T) { - mgr := GeneralManager{} - mgr.setFeatureFlag(featureFlagWGProxy) - mgr.setFeatureFlag(featureFlagDnsForwarder) - - err := mgr.unsetFeatureFlag(featureFlagWGProxy) - if err != nil { - t.Errorf("unexpected error: %s", err) - } - if mgr.featureFlags != 2 { - t.Errorf("invalid feature state, expected: %d, got: %d", 2, mgr.featureFlags) - } - - err = mgr.unsetFeatureFlag(featureFlagDnsForwarder) - if err != nil { - t.Errorf("unexpected error: %s", err) - } - if mgr.featureFlags != 0 { - t.Errorf("invalid feature state, expected: %d, got: %d", 0, mgr.featureFlags) - } -} diff --git a/client/internal/ebpf/ebpf/src/dns_fwd.c b/client/internal/ebpf/ebpf/src/dns_fwd.c deleted file mode 100644 index 9f8de2001..000000000 --- a/client/internal/ebpf/ebpf/src/dns_fwd.c +++ /dev/null @@ -1,67 +0,0 @@ -const __u32 map_key_dns_ip = 0; -const __u32 map_key_dns_port = 1; - -struct bpf_map_def SEC("maps") nb_map_dns_ip = { - .type = BPF_MAP_TYPE_ARRAY, - .key_size = sizeof(__u32), - .value_size = sizeof(__u32), - .max_entries = 10, -}; - -struct bpf_map_def SEC("maps") nb_map_dns_port = { - .type = BPF_MAP_TYPE_ARRAY, - .key_size = sizeof(__u32), - .value_size = sizeof(__u16), - .max_entries = 10, -}; - -__be32 dns_ip = 0; -__be16 dns_port = 0; - -// 13568 is 53 in big endian -__be16 GENERAL_DNS_PORT = 13568; - -bool read_settings() { - __u16 *port_value; - __u32 *ip_value; - - // read dns ip - ip_value = bpf_map_lookup_elem(&nb_map_dns_ip, &map_key_dns_ip); - if(!ip_value) { - return false; - } - dns_ip = htonl(*ip_value); - - // read dns port - port_value = bpf_map_lookup_elem(&nb_map_dns_port, &map_key_dns_port); - if (!port_value) { - return false; - } - dns_port = htons(*port_value); - return true; -} - -int xdp_dns_fwd(struct iphdr *ip, struct udphdr *udp) { - if (dns_port == 0) { - if(!read_settings()){ - return XDP_PASS; - } - // bpf_printk("dns port: %d", ntohs(dns_port)); - // bpf_printk("dns ip: %d", ntohl(dns_ip)); - } - - if (udp->dest == GENERAL_DNS_PORT && ip->daddr == dns_ip) { - udp->dest = dns_port; - // Clear the now-stale checksum; zero means "not computed" for IPv4. - udp->check = 0; - return XDP_PASS; - } - - if (udp->source == dns_port && ip->saddr == dns_ip) { - udp->source = GENERAL_DNS_PORT; - udp->check = 0; - return XDP_PASS; - } - - return XDP_PASS; -} diff --git a/client/internal/ebpf/ebpf/src/prog.c b/client/internal/ebpf/ebpf/src/prog.c deleted file mode 100644 index f32103f28..000000000 --- a/client/internal/ebpf/ebpf/src/prog.c +++ /dev/null @@ -1,60 +0,0 @@ -#include -#include // ETH_P_IP -#include -#include -#include -#include -#include -#include "dns_fwd.c" -#include "wg_proxy.c" - -const __u16 flag_feature_wg_proxy = 0b01; -const __u16 flag_feature_dns_fwd = 0b10; - -const __u32 map_key_features = 0; -struct bpf_map_def SEC("maps") nb_features = { - .type = BPF_MAP_TYPE_ARRAY, - .key_size = sizeof(__u32), - .value_size = sizeof(__u16), - .max_entries = 10, -}; - -SEC("xdp") -int nb_xdp_prog(struct xdp_md *ctx) { - __u16 *features; - features = bpf_map_lookup_elem(&nb_features, &map_key_features); - if (!features) { - return XDP_PASS; - } - - void *data = (void *)(long)ctx->data; - void *data_end = (void *)(long)ctx->data_end; - struct ethhdr *eth = data; - struct iphdr *ip = (data + sizeof(struct ethhdr)); - struct udphdr *udp = (data + sizeof(struct ethhdr) + sizeof(struct iphdr)); - - // return early if not enough data - if (data + sizeof(struct ethhdr) + sizeof(struct iphdr) + sizeof(struct udphdr) > data_end){ - return XDP_PASS; - } - - // skip non IPv4 packages - if (eth->h_proto != htons(ETH_P_IP)) { - return XDP_PASS; - } - - // skip non UPD packages - if (ip->protocol != IPPROTO_UDP) { - return XDP_PASS; - } - - if (*features & flag_feature_dns_fwd) { - xdp_dns_fwd(ip, udp); - } - - if (*features & flag_feature_wg_proxy) { - xdp_wg_proxy(ip, udp); - } - return XDP_PASS; -} -char _license[] SEC("license") = "GPL"; diff --git a/client/internal/ebpf/ebpf/src/readme.md b/client/internal/ebpf/ebpf/src/readme.md deleted file mode 100644 index 0ab393dd4..000000000 --- a/client/internal/ebpf/ebpf/src/readme.md +++ /dev/null @@ -1,17 +0,0 @@ -# DNS forwarder - -The agent attach the XDP program to the lo device. We can not use fake address in eBPF because the -traffic does not appear in the eBPF program. The program capture the traffic on wg_ip:53 and -overwrite in it the destination port to 5053. - -# Debug - -The CONFIG_BPF_EVENTS kernel module is required for bpf_printk. -Apply this code to use bpf_printk -``` -#define bpf_printk(fmt, ...) \ - ({ \ - char ____fmt[] = fmt; \ - bpf_trace_printk(____fmt, sizeof(____fmt), ##__VA_ARGS__); \ - }) -``` diff --git a/client/internal/ebpf/ebpf/src/wg_proxy.c b/client/internal/ebpf/ebpf/src/wg_proxy.c deleted file mode 100644 index 5e7474928..000000000 --- a/client/internal/ebpf/ebpf/src/wg_proxy.c +++ /dev/null @@ -1,60 +0,0 @@ -const __u32 map_key_proxy_port = 0; -const __u32 map_key_wg_port = 1; - -struct bpf_map_def SEC("maps") nb_wg_proxy_settings_map = { - .type = BPF_MAP_TYPE_ARRAY, - .key_size = sizeof(__u32), - .value_size = sizeof(__u16), - .max_entries = 10, -}; - -__u16 proxy_port = 0; -__u16 wg_port = 0; - -bool read_port_settings() { - __u16 *value; - value = bpf_map_lookup_elem(&nb_wg_proxy_settings_map, &map_key_proxy_port); - if (!value) { - return false; - } - - proxy_port = *value; - - value = bpf_map_lookup_elem(&nb_wg_proxy_settings_map, &map_key_wg_port); - if (!value) { - return false; - } - wg_port = htons(*value); - - return true; -} - -int xdp_wg_proxy(struct iphdr *ip, struct udphdr *udp) { - if (proxy_port == 0 || wg_port == 0) { - if (!read_port_settings()){ - return XDP_PASS; - } - // bpf_printk("proxy port: %d, wg port: %d", proxy_port, wg_port); - } - - // 2130706433 = 127.0.0.1 - if (ip->daddr != htonl(2130706433)) { - return XDP_PASS; - } - - if (udp->source != wg_port){ - return XDP_PASS; - } - - __be16 new_src_port = udp->dest; - __be16 new_dst_port = htons(proxy_port); - udp->dest = new_dst_port; - udp->source = new_src_port; - - // The ports are covered by the UDP checksum. This is an IPv4 loopback hop - // and the payload is already integrity-protected, so clear the checksum (a - // zero UDP checksum means "not computed" for IPv4) rather than leave a - // stale value the kernel would drop as UDP_CSUM. - udp->check = 0; - return XDP_PASS; -} diff --git a/client/internal/ebpf/ebpf/wg_proxy_linux.go b/client/internal/ebpf/ebpf/wg_proxy_linux.go deleted file mode 100644 index 4e0df7329..000000000 --- a/client/internal/ebpf/ebpf/wg_proxy_linux.go +++ /dev/null @@ -1,41 +0,0 @@ -package ebpf - -import log "github.com/sirupsen/logrus" - -const ( - mapKeyProxyPort uint32 = 0 - mapKeyWgPort uint32 = 1 -) - -func (tf *GeneralManager) LoadWgProxy(proxyPort, wgPort int) error { - log.Debugf("load ebpf WG proxy") - tf.lock.Lock() - defer tf.lock.Unlock() - - err := tf.loadXdp() - if err != nil { - return err - } - - err = tf.bpfObjs.NbWgProxySettingsMap.Put(mapKeyProxyPort, uint16(proxyPort)) - if err != nil { - return err - } - - err = tf.bpfObjs.NbWgProxySettingsMap.Put(mapKeyWgPort, uint16(wgPort)) - if err != nil { - return err - } - - tf.setFeatureFlag(featureFlagWGProxy) - err = tf.bpfObjs.NbFeatures.Put(mapKeyFeatures, tf.featureFlags) - if err != nil { - return err - } - return nil -} - -func (tf *GeneralManager) FreeWGProxy() error { - log.Debugf("free ebpf WG proxy") - return tf.unsetFeatureFlag(featureFlagWGProxy) -} diff --git a/client/internal/ebpf/instantiater_linux.go b/client/internal/ebpf/instantiater_linux.go deleted file mode 100644 index 20d8145b4..000000000 --- a/client/internal/ebpf/instantiater_linux.go +++ /dev/null @@ -1,15 +0,0 @@ -//go:build !android - -package ebpf - -import ( - "github.com/netbirdio/netbird/client/internal/ebpf/ebpf" - "github.com/netbirdio/netbird/client/internal/ebpf/manager" -) - -// GetEbpfManagerInstance is a wrapper function. This encapsulation is required because if the code import the internal -// ebpf package the Go compiler will include the object files. But it is not supported on Android. It can cause instant -// panic on older Android version. -func GetEbpfManagerInstance() manager.Manager { - return ebpf.GetEbpfManagerInstance() -} diff --git a/client/internal/ebpf/instantiater_nonlinux.go b/client/internal/ebpf/instantiater_nonlinux.go deleted file mode 100644 index b7c38733a..000000000 --- a/client/internal/ebpf/instantiater_nonlinux.go +++ /dev/null @@ -1,10 +0,0 @@ -//go:build !linux || android - -package ebpf - -import "github.com/netbirdio/netbird/client/internal/ebpf/manager" - -// GetEbpfManagerInstance return error because ebpf is not supported on all os -func GetEbpfManagerInstance() manager.Manager { - panic("unsupported os") -} diff --git a/client/internal/ebpf/manager/manager.go b/client/internal/ebpf/manager/manager.go deleted file mode 100644 index 25a767090..000000000 --- a/client/internal/ebpf/manager/manager.go +++ /dev/null @@ -1,11 +0,0 @@ -package manager - -import "net/netip" - -// Manager is used to load multiple eBPF programs. E.g., current DNS programs and WireGuard proxy -type Manager interface { - LoadDNSFwd(ip netip.Addr, dnsPort int) error - FreeDNSFwd() error - LoadWgProxy(proxyPort, wgPort int) error - FreeWGProxy() error -} diff --git a/client/internal/elevate/trusted.go b/client/internal/elevate/trusted.go index c11054c45..98e05fde5 100644 --- a/client/internal/elevate/trusted.go +++ b/client/internal/elevate/trusted.go @@ -6,6 +6,17 @@ import ( "path/filepath" ) +// CheckOnlyOwnerWritable reports an error unless path, and every directory +// leading to it, is owned by an account that can already act with the privileges +// the caller holds, and is writable by nobody else. +// +// Exported for callers outside elevation that read a file while privileged and +// then act on what it says: the same question this package asks of an +// executable, asked of a configuration file. +func CheckOnlyOwnerWritable(path string) error { + return checkOnlyOwnerWritable(path) +} + // trustedSelf returns the path of this executable, provided it is one we are // willing to have run as root. // diff --git a/client/internal/engine.go b/client/internal/engine.go index f8b65f7d8..ccff974d8 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -14,12 +14,14 @@ import ( "sort" "strings" "sync" + "sync/atomic" "time" "github.com/hashicorp/go-multierror" "github.com/pion/ice/v4" "github.com/pion/stun/v3" log "github.com/sirupsen/logrus" + wgdevice "golang.zx2c4.com/wireguard/device" "golang.zx2c4.com/wireguard/tun/netstack" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" @@ -40,7 +42,6 @@ import ( dnsconfig "github.com/netbirdio/netbird/client/internal/dns/config" "github.com/netbirdio/netbird/client/internal/dnsfwd" "github.com/netbirdio/netbird/client/internal/expose" - "github.com/netbirdio/netbird/client/internal/ingressgw" "github.com/netbirdio/netbird/client/internal/lazyconn" "github.com/netbirdio/netbird/client/internal/metrics" "github.com/netbirdio/netbird/client/internal/netflow" @@ -56,6 +57,7 @@ import ( "github.com/netbirdio/netbird/client/internal/rosenpass" "github.com/netbirdio/netbird/client/internal/routemanager" "github.com/netbirdio/netbird/client/internal/statemanager" + "github.com/netbirdio/netbird/client/internal/stdnet" "github.com/netbirdio/netbird/client/internal/syncstore" "github.com/netbirdio/netbird/client/internal/updater" "github.com/netbirdio/netbird/client/jobexec" @@ -236,8 +238,17 @@ type Engine struct { wgInterface WGIface + // wgDevice is a lock-free handle on the WireGuard device behind + // wgInterface. Reaching the device through wgInterface requires + // syncMsgMux, which handleSync holds while it adds and removes peers; + // SetPerformance must stay reachable exactly when that work is stuck. + wgDevice atomic.Pointer[wgdevice.Device] + udpMux *udpmux.UniversalUDPMuxDefault + // wgDetector is shared by every ICE agent through the ICE config. + wgDetector *stdnet.WGDetector + // networkSerial is the latest CurrentSerial (state ID) of the network sent by the Management service networkSerial uint64 @@ -254,11 +265,10 @@ type Engine struct { statusRecorder *peer.Status - firewall firewallManager.Manager - routeManager routemanager.Manager - acl acl.Manager - dnsForwardMgr *dnsfwd.Manager - ingressGatewayMgr *ingressgw.Manager + firewall firewallManager.Manager + routeManager routemanager.Manager + acl acl.Manager + dnsForwardMgr *dnsfwd.Manager dnsServer dns.Server @@ -356,6 +366,7 @@ func NewEngine( mgmClient: services.MgmClient, relayManager: services.RelayManager, peerStore: peerstore.NewConnStore(), + wgDetector: stdnet.NewWGDetector(), syncMsgMux: &sync.Mutex{}, config: config, mobileDep: mobileDep, @@ -440,13 +451,6 @@ func (e *Engine) stopLocked() { e.cleanupSSHConfig() - if e.ingressGatewayMgr != nil { - if err := e.ingressGatewayMgr.Close(); err != nil { - log.Warnf("failed to cleanup forward rules: %v", err) - } - e.ingressGatewayMgr = nil - } - if e.srWatcher != nil { e.srWatcher.Close() } @@ -456,7 +460,7 @@ func (e *Engine) stopLocked() { } if e.updateManager != nil { - e.updateManager.SetDownloadOnly() + e.updateManager.ResetMode() } log.Info("cleaning up status recorder states") @@ -651,10 +655,7 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL) log.Errorf("failed to pull up wgInterface [%s]: %s", e.wgInterface.Name(), err.Error()) return fmt.Errorf("up wg interface: %w", err) } - - // Set up notrack rules immediately after proxy is listening to prevent - // conntrack entries from being created before the rules are in place - e.setupWGProxyNoTrack() + e.wgDevice.Store(e.wgInterface.GetWGDevice()) // Start after interface is up since port may have been resolved from 0 or changed if occupied e.shutdownWg.Add(1) @@ -793,23 +794,6 @@ func (e *Engine) initFirewall() error { return nil } -// setupWGProxyNoTrack configures connection tracking exclusion for WireGuard proxy traffic. -// This prevents conntrack/MASQUERADE from affecting loopback traffic between WireGuard and the eBPF proxy. -func (e *Engine) setupWGProxyNoTrack() { - if e.firewall == nil { - return - } - - proxyPort := e.wgInterface.GetProxyPort() - if proxyPort == 0 { - return - } - - if err := e.firewall.SetupEBPFProxyNoTrack(proxyPort, uint16(e.config.WgPort)); err != nil { - log.Warnf("failed to setup ebpf proxy notrack: %v", err) - } -} - func (e *Engine) blockLanAccess() { if e.config.BlockInbound { // no need to set up extra deny rules if inbound is already blocked in general @@ -972,11 +956,13 @@ func (e *Engine) handleAutoUpdateVersion(autoUpdateSettings *mgmProto.AutoUpdate } if autoUpdateSettings == nil { + log.Infof("no auto-update settings received, defaulting to download-only") + e.updateManager.SetDownloadOnly() return } if autoUpdateSettings.Version == disableAutoUpdate { - log.Infof("auto-update is disabled") + log.Infof("auto-update is disabled, switching to download-only") e.updateManager.SetDownloadOnly() return } @@ -1052,7 +1038,11 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error { // back to empty if the FQDN doesn't have the expected shape. dnsName = extractDNSDomainFromFQDN(pc.GetFqdn()) } - result, err := nbnetworkmap.EnvelopeToNetworkMap(e.ctx, envelope, localKey, dnsName) + // With the firewall disabled there is no ACL manager to program, so + // RoutesFirewallRules would be built and then dropped. On a peer that + // routes many network resources that is the single most expensive + // step of the sync. + result, err := nbnetworkmap.EnvelopeToNetworkMap(e.ctx, envelope, localKey, dnsName, e.config.DisableFirewall) if err != nil { return fmt.Errorf("decode network map envelope: %w", err) } @@ -1635,13 +1625,6 @@ func (e *Engine) updateNetworkMap(networkMap *mgmProto.NetworkMap) error { e.updateDNSForwarder(dnsRouteFeatureFlag, fwdEntries) done() - // Ingress forward rules - done = e.phase("forward_rules") - if _, err := e.updateForwardRules(networkMap.GetForwardingRules()); err != nil { - log.Errorf("failed to update forward rules, err: %v", err) - } - done() - log.Debugf("got peers update from Management Service, total peers to connect to = %d", len(networkMap.GetRemotePeers())) done = e.phase("offline_peers") @@ -2144,6 +2127,10 @@ func (e *Engine) close() { log.Debugf("removing Netbird interface %s", e.config.WgIfaceName) if e.wgInterface != nil { + // Drop the handle before the close starts: a retune that loads it + // afterwards would touch a device on its way out and report success + // for an engine that is already gone. + e.wgDevice.Store(nil) if err := e.wgInterface.Close(); err != nil { log.Errorf("failed closing Netbird interface %s %v", e.config.WgIfaceName, err) } @@ -2182,10 +2169,7 @@ func (e *Engine) close() { } func (e *Engine) newWgIface() (*iface.WGIface, error) { - transportNet, err := e.newStdNet() - if err != nil { - log.Errorf("failed to create pion's stdnet: %s", err) - } + transportNet := e.newStdNet() opts := iface.WGIFaceOpts{ IFaceName: e.config.WgIfaceName, @@ -2303,15 +2287,16 @@ type Performance struct { } // SetPerformance applies the given tuning to this engine's live Device. +// +// It deliberately does not take syncMsgMux. Raising the buffer pool cap is the +// recovery path for a device whose pool is exhausted, and an exhausted pool +// blocks peer removal inside handleSync, which holds syncMsgMux for as long as +// it stays blocked. Taking the lock here would make the retune unreachable in +// the one situation that needs it. func (e *Engine) SetPerformance(t Performance) error { - e.syncMsgMux.Lock() - defer e.syncMsgMux.Unlock() - if e.wgInterface == nil { - return fmt.Errorf("wg interface not initialized") - } - dev := e.wgInterface.GetWGDevice() + dev := e.wgDevice.Load() if dev == nil { - return fmt.Errorf("wg device not initialized") + return errors.New("wg device not initialized") } if t.PreallocatedBuffersPerPool != nil { dev.SetPreallocatedBuffersPerPool(*t.PreallocatedBuffersPerPool) @@ -2739,74 +2724,6 @@ func (e *Engine) setForwarderCapture(pc device.PacketCapture) { } } -func (e *Engine) updateForwardRules(rules []*mgmProto.ForwardingRule) ([]firewallManager.ForwardRule, error) { - if e.firewall == nil { - log.Warn("firewall is disabled, not updating forwarding rules") - return nil, nil - } - - if len(rules) == 0 { - if e.ingressGatewayMgr == nil { - return nil, nil - } - - err := e.ingressGatewayMgr.Close() - e.ingressGatewayMgr = nil - e.statusRecorder.SetIngressGwMgr(nil) - return nil, err - } - - if e.ingressGatewayMgr == nil { - mgr := ingressgw.NewManager(e.firewall) - e.ingressGatewayMgr = mgr - e.statusRecorder.SetIngressGwMgr(mgr) - } - - var merr *multierror.Error - forwardingRules := make([]firewallManager.ForwardRule, 0, len(rules)) - for _, rule := range rules { - proto, err := acl.ConvertToFirewallProtocol(rule.GetProtocol()) - if err != nil { - merr = multierror.Append(merr, fmt.Errorf("failed to convert protocol '%s': %w", rule.GetProtocol(), err)) - continue - } - - dstPortInfo, err := convertPortInfo(rule.GetDestinationPort()) - if err != nil { - merr = multierror.Append(merr, fmt.Errorf("invalid destination port '%v': %w", rule.GetDestinationPort(), err)) - continue - } - - translateIP, err := convertToIP(rule.GetTranslatedAddress()) - if err != nil { - merr = multierror.Append(merr, fmt.Errorf("failed to convert translated address '%s': %w", rule.GetTranslatedAddress(), err)) - continue - } - - translatePort, err := convertPortInfo(rule.GetTranslatedPort()) - if err != nil { - merr = multierror.Append(merr, fmt.Errorf("invalid translate port '%v': %w", rule.GetTranslatedPort(), err)) - continue - } - - forwardRule := firewallManager.ForwardRule{ - Protocol: proto, - DestinationPort: *dstPortInfo, - TranslatedAddress: translateIP, - TranslatedPort: *translatePort, - } - - forwardingRules = append(forwardingRules, forwardRule) - } - - log.Infof("updating forwarding rules: %d", len(forwardingRules)) - if err := e.ingressGatewayMgr.Update(forwardingRules); err != nil { - log.Errorf("failed to update forwarding rules: %v", err) - } - - return forwardingRules, nberrors.FormatErrorOrNil(merr) -} - // toExcludedLazyPeers returns the peers that must have an always-active // connection: those that are not lazy by policy (the per-peer lazy state or the // account flag, subject to the local override). diff --git a/client/internal/engine_generic.go b/client/internal/engine_generic.go index 34a75e45b..e293f1636 100644 --- a/client/internal/engine_generic.go +++ b/client/internal/engine_generic.go @@ -15,5 +15,6 @@ func (e *Engine) createICEConfig() icemaker.Config { UDPMux: e.udpMux.SingleSocketUDPMux, UDPMuxSrflx: e.udpMux, NATExternalIPs: e.parseNATExternalIPMappings(), + WGDetector: e.wgDetector, } } diff --git a/client/internal/engine_js.go b/client/internal/engine_js.go index dce3c57fb..0243b7b37 100644 --- a/client/internal/engine_js.go +++ b/client/internal/engine_js.go @@ -13,6 +13,7 @@ func (e *Engine) createICEConfig() icemaker.Config { InterfaceBlackList: e.config.IFaceBlackList, DisableIPv6Discovery: e.config.DisableIPv6Discovery, NATExternalIPs: e.parseNATExternalIPMappings(), + WGDetector: e.wgDetector, } return cfg } diff --git a/client/internal/engine_privileged_test.go b/client/internal/engine_privileged_test.go index 1b047e017..4449b5788 100644 --- a/client/internal/engine_privileged_test.go +++ b/client/internal/engine_privileged_test.go @@ -12,12 +12,12 @@ import ( "testing" "time" - "go.uber.org/mock/gomock" "github.com/google/uuid" log "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.opentelemetry.io/otel" + "go.uber.org/mock/gomock" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" "google.golang.org/grpc" "google.golang.org/grpc/keepalive" @@ -27,6 +27,7 @@ import ( "github.com/netbirdio/netbird/client/iface/wgaddr" "github.com/netbirdio/netbird/client/internal/dns" "github.com/netbirdio/netbird/client/internal/peer" + "github.com/netbirdio/netbird/client/internal/profilemanager" nbssh "github.com/netbirdio/netbird/client/ssh" "github.com/netbirdio/netbird/client/system" nbdns "github.com/netbirdio/netbird/dns" @@ -41,7 +42,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" @@ -81,6 +81,7 @@ func TestEngine_SSH(t *testing.T) { WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"), WgPrivateKey: key, WgPort: 33100, + IFaceBlackList: profilemanager.DefaultInterfaceBlacklist, ServerSSHAllowed: true, MTU: iface.DefaultMTU, SSHKey: sshKey, @@ -204,11 +205,12 @@ func TestEngine_Sync(t *testing.T) { } relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU) engine := NewEngine(ctx, cancel, &EngineConfig{ - WgIfaceName: "utun103", - WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"), - WgPrivateKey: key, - WgPort: 33100, - MTU: iface.DefaultMTU, + WgIfaceName: "utun103", + WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"), + WgPrivateKey: key, + WgPort: 33100, + IFaceBlackList: profilemanager.DefaultInterfaceBlacklist, + MTU: iface.DefaultMTU, }, EngineServices{ SignalClient: &signal.MockClient{}, MgmClient: &mgmt.MockClient{SyncFunc: syncFunc}, @@ -412,11 +414,12 @@ func createEngine(ctx context.Context, cancel context.CancelFunc, setupKey strin wgPort := 33100 + i conf := &EngineConfig{ - WgIfaceName: ifaceName, - WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address), - WgPrivateKey: key, - WgPort: wgPort, - MTU: iface.DefaultMTU, + WgIfaceName: ifaceName, + WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address), + WgPrivateKey: key, + WgPort: wgPort, + IFaceBlackList: profilemanager.DefaultInterfaceBlacklist, + MTU: iface.DefaultMTU, } relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU) @@ -519,8 +522,8 @@ func startManagement(t *testing.T, dataDir, testFile string) (*grpc.Server, stri updateManager := update_channel.NewPeersUpdateManager(metrics) requestBuffer := server.NewAccountRequestBuffer(context.Background(), store) - networkMapController := controller.NewController(context.Background(), store, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersManager), config, nil) - accountManager, err := server.BuildManager(context.Background(), config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, ia, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore) + networkMapController := controller.NewController(context.Background(), store, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", manager.NewEphemeralManager(store, peersManager), config, nil) + accountManager, err := server.BuildManager(context.Background(), config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, ia, metrics, settingsMockManager, permissionsManager, false, cacheStore) if err != nil { return nil, "", err } diff --git a/client/internal/engine_sessionwatch.go b/client/internal/engine_sessionwatch.go index a46d73f87..05b46a465 100644 --- a/client/internal/engine_sessionwatch.go +++ b/client/internal/engine_sessionwatch.go @@ -1,4 +1,4 @@ -//go:build !js +//go:build !js && !android package internal @@ -7,10 +7,12 @@ import ( "github.com/netbirdio/netbird/client/internal/peer" ) -// newSessionWatcher returns the real SSO session expiry watcher for every -// non-wasm build. The js/wasm build gets a no-op stub from -// engine_sessionwatch_js.go so the sessionwatch package (and its timer -// machinery) never links into the wasm binary. +// newSessionWatcher returns the real SSO session expiry watcher. The js/wasm +// build gets a no-op stub from engine_sessionwatch_js.go so the sessionwatch +// package (and its timer machinery) never links into the wasm binary; the +// android build gets a deadline-only watcher from +// engine_sessionwatch_android.go because the app schedules the warnings +// itself. func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher { return sessionwatch.New(recorder) } diff --git a/client/internal/engine_sessionwatch_android.go b/client/internal/engine_sessionwatch_android.go new file mode 100644 index 000000000..8317f9165 --- /dev/null +++ b/client/internal/engine_sessionwatch_android.go @@ -0,0 +1,12 @@ +//go:build android + +package internal + +import ( + "github.com/netbirdio/netbird/client/internal/auth/sessionwatch" + "github.com/netbirdio/netbird/client/internal/peer" +) + +func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher { + return sessionwatch.NewDeadlineOnly(recorder) +} diff --git a/client/internal/engine_stdnet.go b/client/internal/engine_stdnet.go index 1ebb5779c..d6e2b1aa4 100644 --- a/client/internal/engine_stdnet.go +++ b/client/internal/engine_stdnet.go @@ -6,6 +6,6 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" ) -func (e *Engine) newStdNet() (*stdnet.Net, error) { - return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList) +func (e *Engine) newStdNet() *stdnet.Net { + return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList, e.wgDetector) } diff --git a/client/internal/engine_stdnet_android.go b/client/internal/engine_stdnet_android.go index de3c80bcf..7ef996203 100644 --- a/client/internal/engine_stdnet_android.go +++ b/client/internal/engine_stdnet_android.go @@ -2,6 +2,6 @@ package internal import "github.com/netbirdio/netbird/client/internal/stdnet" -func (e *Engine) newStdNet() (*stdnet.Net, error) { - return stdnet.NewNetWithDiscover(e.clientCtx, e.mobileDep.IFaceDiscover, e.config.IFaceBlackList) +func (e *Engine) newStdNet() *stdnet.Net { + return stdnet.NewNetWithDiscover(e.clientCtx, e.mobileDep.IFaceDiscover, e.config.IFaceBlackList, e.wgDetector) } diff --git a/client/internal/engine_test.go b/client/internal/engine_test.go index ec388ac94..bfde05c04 100644 --- a/client/internal/engine_test.go +++ b/client/internal/engine_test.go @@ -65,7 +65,6 @@ type MockWGIface struct { GetStatsFunc func() (map[string]configurer.WGStats, error) GetInterfaceGUIDStringFunc func() (string, error) GetProxyFunc func() wgproxy.Proxy - GetProxyPortFunc func() uint16 GetNetFunc func() *netstack.Net LastActivitiesFunc func() map[string]monotime.Time } @@ -162,13 +161,6 @@ func (m *MockWGIface) GetProxy() wgproxy.Proxy { return m.GetProxyFunc() } -func (m *MockWGIface) GetProxyPort() uint16 { - if m.GetProxyPortFunc != nil { - return m.GetProxyPortFunc() - } - return 0 -} - func (m *MockWGIface) GetNet() *netstack.Net { return m.GetNetFunc() } @@ -696,10 +688,7 @@ func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) { StatusRecorder: peer.NewRecorder("https://mgm"), }, MobileDependency{}) engine.ctx = ctx - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil) opts := iface.WGIFaceOpts{ IFaceName: wgIfaceName, @@ -904,10 +893,7 @@ func TestEngine_UpdateNetworkMapWithDNSUpdate(t *testing.T) { }, MobileDependency{}) engine.ctx = ctx - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil) opts := iface.WGIFaceOpts{ IFaceName: wgIfaceName, Address: wgaddr.MustParseWGAddress(wgAddr), @@ -1507,3 +1493,21 @@ func TestOverlayAddrsFromAllowedIPs(t *testing.T) { }) } } + +func TestEngine_SyncResponsePersistence(t *testing.T) { + e := &Engine{} + + _, err := e.GetLatestSyncResponse() + require.Error(t, err, "persistence is disabled by default") + + e.SetSyncResponsePersistence(true) + e.persistSyncResponse(&mgmtProto.SyncResponse{NetworkMap: &mgmtProto.NetworkMap{Serial: 7}}) + + got, err := e.GetLatestSyncResponse() + require.NoError(t, err) + assert.Equal(t, uint64(7), got.GetNetworkMap().GetSerial()) + + e.SetSyncResponsePersistence(false) + _, err = e.GetLatestSyncResponse() + require.Error(t, err) +} diff --git a/client/internal/iface_common.go b/client/internal/iface_common.go index 8ffa0b102..d772a3a03 100644 --- a/client/internal/iface_common.go +++ b/client/internal/iface_common.go @@ -28,7 +28,6 @@ type wgIfaceBase interface { Up() (*udpmux.UniversalUDPMuxDefault, error) UpdateAddr(newAddr wgaddr.Address) error GetProxy() wgproxy.Proxy - GetProxyPort() uint16 UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error RemoveEndpointAddress(key string) error RemovePeer(peerKey string) error diff --git a/client/internal/ingressgw/manager.go b/client/internal/ingressgw/manager.go deleted file mode 100644 index 605543d1c..000000000 --- a/client/internal/ingressgw/manager.go +++ /dev/null @@ -1,111 +0,0 @@ -package ingressgw - -import ( - "fmt" - "sync" - - "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" -) - -type DNATFirewall interface { - AddDNATRule(fwdRule firewall.ForwardRule) (firewall.Rule, error) - DeleteDNATRule(rule firewall.Rule) error -} - -type RulePair struct { - firewall.ForwardRule - firewall.Rule -} - -type Manager struct { - dnatFirewall DNATFirewall - - rules map[firewall.RuleID]RulePair - rulesMu sync.Mutex -} - -func NewManager(dnatFirewall DNATFirewall) *Manager { - return &Manager{ - dnatFirewall: dnatFirewall, - rules: make(map[firewall.RuleID]RulePair), - } -} - -func (h *Manager) Update(forwardRules []firewall.ForwardRule) error { - h.rulesMu.Lock() - defer h.rulesMu.Unlock() - - var mErr *multierror.Error - - toDelete := make(map[firewall.RuleID]RulePair, len(h.rules)) - for id, r := range h.rules { - toDelete[id] = r - } - - // Process new/updated rules - for _, fwdRule := range forwardRules { - id := fwdRule.ID() - if _, ok := h.rules[id]; ok { - delete(toDelete, id) - continue - } - - rule, err := h.dnatFirewall.AddDNATRule(fwdRule) - if err != nil { - mErr = multierror.Append(mErr, fmt.Errorf("add forward rule '%s': %v", fwdRule.String(), err)) - continue - } - if rule == nil { - mErr = multierror.Append(mErr, fmt.Errorf("add forward rule '%s': backend returned no rule", fwdRule.String())) - continue - } - log.Infof("forward rule has been added '%s'", fwdRule) - h.rules[id] = RulePair{ - ForwardRule: fwdRule, - Rule: rule, - } - } - - // Remove deleted rules - for id, rulePair := range toDelete { - if err := h.dnatFirewall.DeleteDNATRule(rulePair.Rule); err != nil { - mErr = multierror.Append(mErr, fmt.Errorf("failed to delete forward rule '%s': %v", rulePair.ForwardRule.String(), err)) - } - log.Infof("forward rule has been deleted '%s'", rulePair.ForwardRule) - delete(h.rules, id) - } - - return nberrors.FormatErrorOrNil(mErr) -} - -func (h *Manager) Close() error { - h.rulesMu.Lock() - defer h.rulesMu.Unlock() - - log.Infof("clean up all (%d) forward rules", len(h.rules)) - var mErr *multierror.Error - for _, rule := range h.rules { - if err := h.dnatFirewall.DeleteDNATRule(rule.Rule); err != nil { - mErr = multierror.Append(mErr, fmt.Errorf("failed to delete forward rule '%s': %v", rule, err)) - } - } - - h.rules = make(map[firewall.RuleID]RulePair) - return nberrors.FormatErrorOrNil(mErr) -} - -func (h *Manager) Rules() []firewall.ForwardRule { - h.rulesMu.Lock() - defer h.rulesMu.Unlock() - - rules := make([]firewall.ForwardRule, 0, len(h.rules)) - for _, rulePair := range h.rules { - rules = append(rules, rulePair.ForwardRule) - } - - return rules -} diff --git a/client/internal/ingressgw/manager_test.go b/client/internal/ingressgw/manager_test.go deleted file mode 100644 index 0cd40fcc4..000000000 --- a/client/internal/ingressgw/manager_test.go +++ /dev/null @@ -1,281 +0,0 @@ -package ingressgw - -import ( - "fmt" - "net/netip" - "testing" - - firewall "github.com/netbirdio/netbird/client/firewall/manager" -) - -var ( - _ firewall.Rule = (*MocFwRule)(nil) - _ DNATFirewall = &MockDNATFirewall{} -) - -type MocFwRule struct { - id firewall.RuleID -} - -func (m *MocFwRule) ID() firewall.RuleID { - return m.id -} - -type MockDNATFirewall struct { - throwError bool -} - -func (m *MockDNATFirewall) AddDNATRule(fwdRule firewall.ForwardRule) (firewall.Rule, error) { - if m.throwError { - return nil, fmt.Errorf("moc error") - } - - fwRule := &MocFwRule{ - id: fwdRule.ID(), - } - return fwRule, nil -} - -func (m *MockDNATFirewall) DeleteDNATRule(rule firewall.Rule) error { - if m.throwError { - return fmt.Errorf("moc error") - } - return nil -} - -func (m *MockDNATFirewall) forceToThrowErrors() { - m.throwError = true -} - -func TestManager_AddRule(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - - updates := []firewall.ForwardRule{ - { - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - }, - { - Protocol: firewall.ProtocolUDP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - }} - - if err := mgr.Update(updates); err != nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != len(updates) { - t.Errorf("unexpected rules count: %d", len(rules)) - } -} - -func TestManager_UpdateRule(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - ruleUDP := firewall.ForwardRule{ - Protocol: firewall.ProtocolUDP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.2"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleUDP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != 1 { - t.Errorf("unexpected rules count: %d", len(rules)) - } - - if rules[0].TranslatedAddress.String() != ruleUDP.TranslatedAddress.String() { - t.Errorf("unexpected rule: %v", rules[0]) - } - - if rules[0].TranslatedPort.String() != ruleUDP.TranslatedPort.String() { - t.Errorf("unexpected rule: %v", rules[0]) - } - - if rules[0].DestinationPort.String() != ruleUDP.DestinationPort.String() { - t.Errorf("unexpected rule: %v", rules[0]) - } - - if rules[0].Protocol != ruleUDP.Protocol { - t.Errorf("unexpected rule: %v", rules[0]) - } -} - -func TestManager_ExtendRules(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - ruleUDP := firewall.ForwardRule{ - Protocol: firewall.ProtocolUDP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.2"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP, ruleUDP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != 2 { - t.Errorf("unexpected rules count: %d", len(rules)) - } -} - -func TestManager_UnderlingError(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - ruleUDP := firewall.ForwardRule{ - Protocol: firewall.ProtocolUDP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.2"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - fw.forceToThrowErrors() - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP, ruleUDP}); err == nil { - t.Errorf("expected error") - } - - rules := mgr.Rules() - if len(rules) != 1 { - t.Errorf("unexpected rules count: %d", len(rules)) - } -} - -func TestManager_Cleanup(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - if err := mgr.Update([]firewall.ForwardRule{}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != 0 { - t.Errorf("unexpected rules count: %d", len(rules)) - } -} - -func TestManager_DeleteBrokenRule(t *testing.T) { - fw := &MockDNATFirewall{} - - // force to throw errors when Add DNAT Rule - fw.forceToThrowErrors() - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err == nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != 0 { - t.Errorf("unexpected rules count: %d", len(rules)) - } - - // simulate that to remove a broken rule - if err := mgr.Update([]firewall.ForwardRule{}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - if err := mgr.Close(); err != nil { - t.Errorf("unexpected error: %v", err) - } -} - -func TestManager_Close(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - if err := mgr.Close(); err != nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != 0 { - t.Errorf("unexpected rules count: %d", len(rules)) - } -} diff --git a/client/internal/ipcauth/forward_test.go b/client/internal/ipcauth/forward_test.go index d9adf05da..d80c293be 100644 --- a/client/internal/ipcauth/forward_test.go +++ b/client/internal/ipcauth/forward_test.go @@ -38,7 +38,7 @@ func asDaemon(t *testing.T, id Identity) { prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate }) selfIdentity, selfKnown = id, true - selfMayDelegate = !id.IsPrivileged() + selfMayDelegate = mayDelegate(id) } func TestCallerIdentity_DirectConnections(t *testing.T) { diff --git a/client/internal/ipcauth/identity.go b/client/internal/ipcauth/identity.go index d7d10f57d..255585821 100644 --- a/client/internal/ipcauth/identity.go +++ b/client/internal/ipcauth/identity.go @@ -18,7 +18,8 @@ import ( "google.golang.org/grpc/peer" ) -// Well-known Windows SIDs that identify a fully privileged principal. +// Well-known Windows SIDs. Only LocalSystem and BUILTIN\Administrators identify a +// privileged principal; the service accounts are shared by unrelated services. const ( sidLocalSystem = "S-1-5-18" // NT AUTHORITY\SYSTEM sidLocalService = "S-1-5-19" // NT AUTHORITY\LOCAL SERVICE @@ -67,9 +68,9 @@ func (i Identity) IsWindows() bool { // user-to-root boundary. // // On Windows the decision comes from the caller's token rather than from -// account names or group RIDs: an elevated token, one of the service accounts -// the daemon itself may run as, or a token with BUILTIN\Administrators -// enabled. A UAC-filtered administrator has that group marked deny-only, and +// account names or group RIDs: an elevated token, the LocalSystem SID, or a +// token with BUILTIN\Administrators enabled. LocalService and NetworkService +// are not privileged by SID. A UAC-filtered administrator has that group marked deny-only, and // deny-only groups are dropped when the identity is captured, so such a // caller is correctly reported as unprivileged. Domain group memberships // (Domain Admins and friends) are deliberately not consulted: they say @@ -83,8 +84,7 @@ func (i Identity) IsPrivileged() bool { return true } - switch i.SID { - case sidLocalSystem, sidLocalService, sidNetworkService: + if i.SID == sidLocalSystem { return true } diff --git a/client/internal/ipcauth/identity_sameuser_test.go b/client/internal/ipcauth/identity_test.go similarity index 56% rename from client/internal/ipcauth/identity_sameuser_test.go rename to client/internal/ipcauth/identity_test.go index c98f583db..57be1b94e 100644 --- a/client/internal/ipcauth/identity_sameuser_test.go +++ b/client/internal/ipcauth/identity_test.go @@ -64,3 +64,58 @@ func TestIdentitySameUser(t *testing.T) { }) } } + +func TestIdentityIsPrivileged(t *testing.T) { + tests := []struct { + name string + id Identity + want bool + }{ + { + name: "Root", + id: Identity{UID: 0, GID: 0}, + want: true, + }, + { + name: "Non-root", + id: Identity{UID: 1000, GID: 1000}, + want: false, + }, + { + name: "Local system windows", + id: Identity{SID: sidLocalSystem}, + want: true, + }, + { + name: "Windows elevated", + id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001", Elevated: true}, + want: true, + }, + { + name: "Admin group windows", + id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001", Groups: []string{sidAdministrators}}, + want: true, + }, + { + name: "Regular user windows", + id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001"}, + want: false, + }, + { + name: "Network service windows", + id: Identity{SID: sidNetworkService}, + want: false, + }, + { + name: "Local service windows", + id: Identity{SID: sidLocalService}, + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, tt.id.IsPrivileged()) + }) + } +} diff --git a/client/internal/ipcauth/privileged.go b/client/internal/ipcauth/privileged.go index 3c2e68432..54c66a5d2 100644 --- a/client/internal/ipcauth/privileged.go +++ b/client/internal/ipcauth/privileged.go @@ -45,7 +45,15 @@ func init() { // matching there would let a non-elevated shell of an administrator account // act as an administrator, which is the boundary the token check exists to // keep. - selfMayDelegate = !id.IsPrivileged() + selfMayDelegate = mayDelegate(id) +} + +// mayDelegate reports whether a daemon running as id may extend its authority to +// callers sharing its identity. The shared service accounts are excluded: their +// SID is held by unrelated services, so matching on it would grant them the +// daemon's authority. +func mayDelegate(id Identity) bool { + return !id.IsPrivileged() && id.SID != sidLocalService && id.SID != sidNetworkService } // IsDaemonSelf reports whether an identity is this very process. The JSON gateway diff --git a/client/internal/ipcauth/privileged_test.go b/client/internal/ipcauth/privileged_test.go index c1c7c1543..a6bbcf44b 100644 --- a/client/internal/ipcauth/privileged_test.go +++ b/client/internal/ipcauth/privileged_test.go @@ -98,7 +98,7 @@ func TestIsPrivilegedCaller_SelfRule(t *testing.T) { t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate }) selfIdentity, selfKnown = tt.self, tt.selfKnown - selfMayDelegate = tt.selfKnown && !tt.self.IsPrivileged() + selfMayDelegate = tt.selfKnown && mayDelegate(tt.self) if got := IsPrivilegedCaller(tt.caller); got != tt.want { t.Fatalf("IsPrivilegedCaller(%v) with daemon %v = %t, want %t", @@ -132,3 +132,28 @@ func TestIsPrivilegedCaller_ThisProcess(t *testing.T) { t.Errorf("an unrelated identity %v was treated as privileged", other) } } + +// The shared service accounts are held by unrelated services, so a daemon running +// as one of them must not extend its authority to every process with that SID. +func TestMayDelegate(t *testing.T) { + tests := []struct { + name string + self Identity + want bool + }{ + {name: "unprivileged unix user", self: Identity{UID: 1000}, want: true}, + {name: "root", self: Identity{UID: 0}, want: false}, + {name: "unprivileged windows user", self: Identity{SID: "S-1-5-21-1-2-3-1001"}, want: true}, + {name: "elevated windows user", self: Identity{SID: "S-1-5-21-1-2-3-1001", Elevated: true}, want: false}, + {name: "local system", self: Identity{SID: sidLocalSystem}, want: false}, + {name: "local service", self: Identity{SID: sidLocalService}, want: false}, + {name: "network service", self: Identity{SID: sidNetworkService}, want: false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := mayDelegate(tt.self); got != tt.want { + t.Errorf("mayDelegate(%+v) = %v, want %v", tt.self, got, tt.want) + } + }) + } +} diff --git a/client/internal/message_convert.go b/client/internal/message_convert.go deleted file mode 100644 index 60f19e228..000000000 --- a/client/internal/message_convert.go +++ /dev/null @@ -1,43 +0,0 @@ -package internal - -import ( - "errors" - "fmt" - "net" - "net/netip" - - firewallManager "github.com/netbirdio/netbird/client/firewall/manager" - mgmProto "github.com/netbirdio/netbird/shared/management/proto" -) - -func convertPortInfo(portInfo *mgmProto.PortInfo) (*firewallManager.Port, error) { - if portInfo == nil { - return nil, errors.New("portInfo cannot be nil") - } - - if portInfo.GetPort() != 0 { - return firewallManager.NewPort(int(portInfo.GetPort())) - } - - if portInfo.GetRange() != nil { - return firewallManager.NewPort(int(portInfo.GetRange().Start), int(portInfo.GetRange().End)) - } - - return nil, fmt.Errorf("invalid portInfo: %v", portInfo) -} - -func convertToIP(rawIP []byte) (netip.Addr, error) { - if rawIP == nil { - return netip.Addr{}, errors.New("input bytes cannot be nil") - } - - if len(rawIP) != net.IPv4len && len(rawIP) != net.IPv6len { - return netip.Addr{}, fmt.Errorf("invalid IP length: %d", len(rawIP)) - } - - if len(rawIP) == net.IPv4len { - return netip.AddrFrom4([4]byte(rawIP)), nil - } - - return netip.AddrFrom16([16]byte(rawIP)), nil -} diff --git a/client/internal/peer/conn.go b/client/internal/peer/conn.go index 83089606f..17823e043 100644 --- a/client/internal/peer/conn.go +++ b/client/internal/peer/conn.go @@ -135,9 +135,10 @@ type Conn struct { // used to store the remote Rosenpass key for Relayed connection in case of connection update from ice rosenpassRemoteKey []byte - wgProxyICE wgproxy.Proxy - wgProxyRelay wgproxy.Proxy - handshaker *Handshaker + wgProxyICE wgproxy.Proxy + wgProxyRelay wgproxy.Proxy + relayedConnRef *relayClient.Conn + handshaker *Handshaker guard *guard.Guard wg sync.WaitGroup @@ -560,7 +561,7 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) { conn.mu.Lock() defer conn.mu.Unlock() - if conn.ctx.Err() != nil { + if conn.ctx.Err() != nil || rci.relayedConn.Context().Err() != nil { if err := rci.relayedConn.Close(); err != nil { conn.Log.Warnf("failed to close unnecessary relayed connection: %v", err) } @@ -575,7 +576,9 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) { conn.Log.Errorf("failed to add relayed net.Conn to local proxy: %v", err) return } - wgProxy.SetDisconnectListener(conn.onRelayDisconnected) + wgProxy.SetDisconnectListener(func() { + conn.onRelayDisconnected(rci.relayedConn) + }) conn.dumpState.NewLocalProxy() @@ -583,7 +586,7 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) { if conn.isICEActive() { conn.Log.Debugf("do not switch to relay because current priority is: %s", conn.currentConnPriority.String()) - conn.setRelayedProxy(wgProxy) + conn.setRelayedProxy(wgProxy, rci.relayedConn) conn.statusRelay.SetConnected() conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, time.Now()) return @@ -614,15 +617,26 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) { conn.rosenpassRemoteKey = rci.rosenpassPubKey conn.currentConnPriority = conntype.Relay conn.statusRelay.SetConnected() - conn.setRelayedProxy(wgProxy) + conn.setRelayedProxy(wgProxy, rci.relayedConn) conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, updateTime) conn.Log.Infof("start to communicate with peer via relay") conn.doOnConnected(rci.rosenpassPubKey, rci.rosenpassAddr, updateTime) } -func (conn *Conn) onRelayDisconnected() { +// onRelayDisconnected reports the teardown of a relayed connection. relayedConn +// names the connection the signal belongs to, so a signal that arrives after +// its connection was replaced is ignored instead of tearing down its successor. +// A nil relayedConn means the caller does not track generations and the current +// connection is always torn down. +func (conn *Conn) onRelayDisconnected(relayedConn *relayClient.Conn) { conn.mu.Lock() defer conn.mu.Unlock() + + if relayedConn != nil && conn.relayedConnRef != relayedConn { + conn.Log.Debugf("ignoring relay disconnect of a superseded connection") + return + } + conn.handleRelayDisconnectedLocked() } @@ -646,6 +660,7 @@ func (conn *Conn) handleRelayDisconnectedLocked() { _ = conn.wgProxyRelay.CloseConn() conn.wgProxyRelay = nil } + conn.relayedConnRef = nil changed := conn.statusRelay.Get() != worker.StatusDisconnected if changed { @@ -813,7 +828,8 @@ func (conn *Conn) evalStatus() ConnStatus { // // The result is a tri-state: // - ConnStatusConnected: all available transports are up -// - ConnStatusPartiallyConnected: relay is up but ICE is still pending/reconnecting +// - ConnStatusPartiallyConnected: one transport carries the traffic and the other does +// not: relay up with ICE down, or ICE up with the shared relay transport down // - ConnStatusDisconnected: no working transport func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) { defer func() { @@ -830,13 +846,14 @@ func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) { } return evalConnStatus(connStatusInputs{ - forceRelay: IsForceRelayed(), - peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(), - relayConnected: conn.statusRelay.Get() == worker.StatusConnected, - remoteSupportsICE: conn.handshaker.RemoteICESupported(), - iceWorkerCreated: iceWorkerCreated, - iceStatusConnecting: conn.statusICE.Get() != worker.StatusDisconnected, - iceInProgress: iceInProgress, + forceRelay: IsForceRelayed(), + peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(), + relayConnected: conn.statusRelay.Get() == worker.StatusConnected, + relayTransportConnected: conn.workerRelay.IsTransportConnected(), + remoteSupportsICE: conn.handshaker.RemoteICESupported(), + iceWorkerCreated: iceWorkerCreated, + iceStatusConnected: conn.statusICE.Get() == worker.StatusConnected, + iceInProgress: iceInProgress, }) } @@ -930,13 +947,14 @@ func (conn *Conn) logTraceConnState() { } } -func (conn *Conn) setRelayedProxy(proxy wgproxy.Proxy) { +func (conn *Conn) setRelayedProxy(proxy wgproxy.Proxy, relayedConn *relayClient.Conn) { if conn.wgProxyRelay != nil { if err := conn.wgProxyRelay.CloseConn(); err != nil { conn.Log.Warnf("failed to close deprecated wg proxy conn: %v", err) } } conn.wgProxyRelay = proxy + conn.relayedConnRef = relayedConn } // onWGHandshakeSuccess is called when the first WireGuard handshake is detected @@ -1044,19 +1062,21 @@ func evalConnStatus(in connStatusInputs) guard.ConnStatus { return boolToConnStatus(relayUsedAndUp) } - // ICE counts as "up" when the status is anything other than Disconnected, OR - // when a negotiation is currently in progress (so we don't spam offers while one is in flight). - iceUp := in.iceStatusConnecting || in.iceInProgress + // ICE counts as "running" when either connected or attempting to connect. + iceRunning := in.iceStatusConnected || in.iceInProgress // Relay side is acceptable if the peer doesn't rely on relay, or relay is connected. relayOK := !in.peerUsesRelay || in.relayConnected switch { - case iceUp && relayOK: + case iceRunning && relayOK: return guard.ConnStatusConnected case relayUsedAndUp: // Relay is up but ICE is down — partially connected. return guard.ConnStatusPartiallyConnected + case in.iceStatusConnected && !in.relayTransportConnected: + // ICE is up and the shared relay transport is down — offers cannot restore it. + return guard.ConnStatusPartiallyConnected default: return guard.ConnStatusDisconnected } diff --git a/client/internal/peer/conn_status.go b/client/internal/peer/conn_status.go index d6ad37b70..acf271534 100644 --- a/client/internal/peer/conn_status.go +++ b/client/internal/peer/conn_status.go @@ -17,13 +17,14 @@ const ( // tri-state connection classification. Extracted so the decision logic can be unit-tested // without constructing full Worker/Handshaker objects. type connStatusInputs struct { - forceRelay bool // NB_FORCE_RELAY or JS/WASM - peerUsesRelay bool // remote peer advertises relay support AND local has relay - relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay) - remoteSupportsICE bool // remote peer sent ICE credentials - iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode) - iceStatusConnecting bool // statusICE is anything other than Disconnected - iceInProgress bool // a negotiation is currently in flight + forceRelay bool // NB_FORCE_RELAY or JS/WASM + peerUsesRelay bool // remote peer advertises relay support AND local has relay + relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay) + relayTransportConnected bool // the relay transport shared by all peers on that server is up + remoteSupportsICE bool // remote peer sent ICE credentials + iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode) + iceStatusConnected bool // statusICE reports Connected + iceInProgress bool // a negotiation is currently in flight } // ConnStatus describe the status of a peer's connection diff --git a/client/internal/peer/conn_status_eval_test.go b/client/internal/peer/conn_status_eval_test.go index 66393cafe..a239196dc 100644 --- a/client/internal/peer/conn_status_eval_test.go +++ b/client/internal/peer/conn_status_eval_test.go @@ -30,6 +30,21 @@ func TestEvalConnStatus_ForceRelay(t *testing.T) { }, want: guard.ConnStatusDisconnected, }, + { + name: "force relay, relay up but the shared transport reports down", + in: connStatusInputs{ + forceRelay: true, + peerUsesRelay: true, + relayConnected: true, + relayTransportConnected: false, + // The ICE inputs are set so that the force-relay return is the only branch + // that can produce Connected here: without it the peer would fall through to + // relayUsedAndUp and report PartiallyConnected. + remoteSupportsICE: true, + iceWorkerCreated: true, + }, + want: guard.ConnStatusConnected, + }, { name: "force relay, peer does NOT use relay - disconnected forever", in: connStatusInputs{ @@ -123,24 +138,28 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = true in.relayConnected = true - in.iceStatusConnecting = true + in.relayTransportConnected = true + in.iceStatusConnected = true }, want: guard.ConnStatusConnected, }, { - name: "ICE connected, peer does NOT use relay", + name: "ICE connected, peer does NOT use relay, shared transport down", mutator: func(in *connStatusInputs) { in.peerUsesRelay = false in.relayConnected = false - in.iceStatusConnecting = true + in.relayTransportConnected = false + in.iceStatusConnected = true }, + // A peer that does not rely on relay is unaffected by the shared transport: + // relayOK is true, so the first arm matches before the transport is considered. want: guard.ConnStatusConnected, }, { name: "ICE InProgress only, peer does NOT use relay", mutator: func(in *connStatusInputs) { in.peerUsesRelay = false - in.iceStatusConnecting = false + in.iceStatusConnected = false in.iceInProgress = true }, want: guard.ConnStatusConnected, @@ -150,7 +169,8 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = true in.relayConnected = true - in.iceStatusConnecting = false + in.relayTransportConnected = true + in.iceStatusConnected = false in.iceInProgress = false }, want: guard.ConnStatusPartiallyConnected, @@ -160,21 +180,60 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = false in.relayConnected = false - in.iceStatusConnecting = false + in.iceStatusConnected = false in.iceInProgress = false }, want: guard.ConnStatusDisconnected, }, { - name: "ICE up, peer uses relay but relay down -> partial (relay required, ICE ignored)", + name: "ICE connected, relay down for this peer but the shared transport is up -> disconnected", mutator: func(in *connStatusInputs) { in.peerUsesRelay = true in.relayConnected = false - in.iceStatusConnecting = true + in.relayTransportConnected = true + in.iceStatusConnected = true + }, + // The transport is fine, so the peer itself is unreachable over relay: it may have + // moved to another server, and only an offer carries its new relay address. + want: guard.ConnStatusDisconnected, + }, + { + name: "ICE connected, the shared relay transport is down -> partial", + mutator: func(in *connStatusInputs) { + in.peerUsesRelay = true + in.relayConnected = false + in.relayTransportConnected = false + in.iceStatusConnected = true + }, + // ICE carries the traffic and the relay transport is restored by the relay client's + // own guard, not by offers, so this must not trigger the aggressive retry. + want: guard.ConnStatusPartiallyConnected, + }, + { + name: "ICE only negotiating while the shared relay transport is down -> disconnected", + mutator: func(in *connStatusInputs) { + in.peerUsesRelay = true + in.relayConnected = false + in.relayTransportConnected = false + in.iceStatusConnected = false + in.iceInProgress = true + }, + // A negotiation in flight is not a working transport, so this peer has no path at + // all and must keep the aggressive retry. Calling it partially connected spends the + // ICE retry budget and parks the guard on the hourly ticker, and nothing wakes it + // when the negotiation then fails: onICEStateDisconnected is only reached once ICE + // has reached Connected (worker_ice.go onConnectionStateChange). + want: guard.ConnStatusDisconnected, + }, + { + name: "ICE down and the shared relay transport is down -> disconnected", + mutator: func(in *connStatusInputs) { + in.peerUsesRelay = true + in.relayConnected = false + in.relayTransportConnected = false + in.iceStatusConnected = false + in.iceInProgress = false }, - // relayOK = false (peer uses relay but it's down), iceUp = true - // first switch arm fails (relayOK false), relayUsedAndUp = false (relay down), - // falls into default: Disconnected. want: guard.ConnStatusDisconnected, }, { @@ -182,7 +241,7 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = false in.relayConnected = true // not actually used since peer doesn't rely on it - in.iceStatusConnecting = false + in.iceStatusConnected = false in.iceInProgress = false }, want: guard.ConnStatusDisconnected, diff --git a/client/internal/peer/conn_test.go b/client/internal/peer/conn_test.go index b709d5e40..8f3d6216d 100644 --- a/client/internal/peer/conn_test.go +++ b/client/internal/peer/conn_test.go @@ -42,7 +42,7 @@ func TestNewConn_interfaceFilter(t *testing.T) { ignore := []string{iface.WgInterfaceDefault, "tun0", "zt", "ZeroTier", "utun", "wg", "ts", "Tailscale", "tailscale"} - filter := stdnet.InterfaceFilter(ignore) + filter := stdnet.InterfaceFilter(ignore, nil) for _, s := range ignore { assert.Equal(t, filter(s), false) diff --git a/client/internal/peer/guard/guard.go b/client/internal/peer/guard/guard.go index 73bab2a89..15028d91c 100644 --- a/client/internal/peer/guard/guard.go +++ b/client/internal/peer/guard/guard.go @@ -14,7 +14,8 @@ type ConnStatus int const ( // ConnStatusDisconnected means neither ICE nor Relay is connected. ConnStatusDisconnected ConnStatus = iota - // ConnStatusPartiallyConnected means Relay is connected but ICE is not. + // ConnStatusPartiallyConnected means one transport is usable and the other is not: + // relay connected with ICE down, or ICE connected with the shared relay transport down. ConnStatusPartiallyConnected // ConnStatusConnected means all required connections are established. ConnStatusConnected @@ -87,8 +88,9 @@ func (g *Guard) SetICEConnDisconnected() { // - Connected: no action, the peer is fully reachable. // - Disconnected (neither ICE nor Relay): retries aggressively with exponential backoff (800ms doubling // up to timeout), never gives up. This ensures rapid recovery when the peer has no connectivity at all. -// - PartiallyConnected (Relay up, ICE not): retries up to 3 times with exponential backoff, then switches -// to one attempt per hour. This limits signaling traffic when relay already provides connectivity. +// - PartiallyConnected (one transport usable, the other not): retries up to 3 times +// with exponential backoff, then switches to one attempt per hour. This limits +// signaling traffic while the peer still has a working path. // // External events (relay/ICE disconnect, signal/relay reconnect, candidate changes) reset the retry // counter and backoff ticker, giving ICE a fresh chance after network conditions change. diff --git a/client/internal/peer/handshaker.go b/client/internal/peer/handshaker.go index 6ecb2a947..654e32158 100644 --- a/client/internal/peer/handshaker.go +++ b/client/internal/peer/handshaker.go @@ -116,7 +116,7 @@ func (h *Handshaker) Listen(ctx context.Context) { for { select { case remoteOfferAnswer := <-h.remoteOffersCh: - h.log.Infof("received offer, running version %s, remote WireGuard listen port %d, session id: %s, remote ICE supported: %t", remoteOfferAnswer.Version, remoteOfferAnswer.WgListenPort, remoteOfferAnswer.SessionIDString(), remoteOfferAnswer.hasICECredentials()) + h.log.Infof("received offer, running version %s, remote WireGuard listen port %d, session id: %s, remote ICE supported: %t, relay server: %s, relay IP: %s", remoteOfferAnswer.Version, remoteOfferAnswer.WgListenPort, remoteOfferAnswer.SessionIDString(), remoteOfferAnswer.hasICECredentials(), remoteOfferAnswer.RelaySrvAddress, remoteOfferAnswer.RelaySrvIP) // Record signaling received for reconnection attempts if h.metricsStages != nil { @@ -138,7 +138,7 @@ func (h *Handshaker) Listen(ctx context.Context) { continue } case remoteOfferAnswer := <-h.remoteAnswerCh: - h.log.Infof("received answer, running version %s, remote WireGuard listen port %d, session id: %s, remote ICE supported: %t", remoteOfferAnswer.Version, remoteOfferAnswer.WgListenPort, remoteOfferAnswer.SessionIDString(), remoteOfferAnswer.hasICECredentials()) + h.log.Infof("received answer, running version %s, remote WireGuard listen port %d, session id: %s, remote ICE supported: %t, relay server: %s, relay IP: %s", remoteOfferAnswer.Version, remoteOfferAnswer.WgListenPort, remoteOfferAnswer.SessionIDString(), remoteOfferAnswer.hasICECredentials(), remoteOfferAnswer.RelaySrvAddress, remoteOfferAnswer.RelaySrvIP) // Record signaling received for reconnection attempts if h.metricsStages != nil { @@ -209,14 +209,14 @@ func (h *Handshaker) sendOffer() error { } offer := h.buildOfferAnswer() - h.log.Debugf("sending offer with serial: %s", offer.SessionIDString()) + h.log.Debugf("sending offer with serial: %s, relay server: %s, relay IP: %s", offer.SessionIDString(), offer.RelaySrvAddress, offer.RelaySrvIP) return h.signaler.SignalOffer(offer, h.config.Key) } func (h *Handshaker) sendAnswer() error { answer := h.buildOfferAnswer() - h.log.Debugf("sending answer with serial: %s", answer.SessionIDString()) + h.log.Debugf("sending answer with serial: %s, relay server: %s, relay IP: %s", answer.SessionIDString(), answer.RelaySrvAddress, answer.RelaySrvIP) return h.signaler.SignalAnswer(answer, h.config.Key) } diff --git a/client/internal/peer/ice/agent.go b/client/internal/peer/ice/agent.go index c74b46d10..a919cded5 100644 --- a/client/internal/peer/ice/agent.go +++ b/client/internal/peer/ice/agent.go @@ -39,10 +39,7 @@ func NewAgent(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, c iceFailedTimeout := iceFailedTimeout() iceRelayAcceptanceMinWait := iceRelayAcceptanceMinWait() - transportNet, err := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList) - if err != nil { - log.Errorf("failed to create pion's stdnet: %s", err) - } + transportNet := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList, config.WGDetector) fac := logging.NewDefaultLoggerFactory() @@ -53,7 +50,7 @@ func NewAgent(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, c NetworkTypes: []ice.NetworkType{ice.NetworkTypeUDP4, ice.NetworkTypeUDP6}, Urls: config.StunTurn.Load(), CandidateTypes: candidateTypes, - InterfaceFilter: stdnet.InterfaceFilter(config.InterfaceBlackList), + InterfaceFilter: stdnet.InterfaceFilter(config.InterfaceBlackList, config.WGDetector), UDPMux: config.UDPMux, UDPMuxSrflx: config.UDPMuxSrflx, NAT1To1IPs: config.NATExternalIPs, diff --git a/client/internal/peer/ice/config.go b/client/internal/peer/ice/config.go index dd5d67403..9d6b19925 100644 --- a/client/internal/peer/ice/config.go +++ b/client/internal/peer/ice/config.go @@ -2,6 +2,8 @@ package ice import ( "github.com/pion/ice/v4" + + "github.com/netbirdio/netbird/client/internal/stdnet" ) type Config struct { @@ -17,4 +19,8 @@ type Config struct { UDPMuxSrflx ice.UniversalUDPMux NATExternalIPs []string + + // WGDetector is shared by every agent so that the WireGuard check the interface + // filter performs is not repeated for each of them. + WGDetector *stdnet.WGDetector } diff --git a/client/internal/peer/ice/stdnet.go b/client/internal/peer/ice/stdnet.go index 685ed0363..fbbe4f388 100644 --- a/client/internal/peer/ice/stdnet.go +++ b/client/internal/peer/ice/stdnet.go @@ -8,6 +8,6 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" ) -func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) { - return stdnet.NewNet(ctx, ifaceBlacklist) +func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string, detector *stdnet.WGDetector) *stdnet.Net { + return stdnet.NewNet(ctx, ifaceBlacklist, detector) } diff --git a/client/internal/peer/ice/stdnet_android.go b/client/internal/peer/ice/stdnet_android.go index 5033ec1b9..72f0c329e 100644 --- a/client/internal/peer/ice/stdnet_android.go +++ b/client/internal/peer/ice/stdnet_android.go @@ -6,6 +6,6 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" ) -func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) { - return stdnet.NewNetWithDiscover(ctx, iFaceDiscover, ifaceBlacklist) +func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string, detector *stdnet.WGDetector) *stdnet.Net { + return stdnet.NewNetWithDiscover(ctx, iFaceDiscover, ifaceBlacklist, detector) } diff --git a/client/internal/peer/notifier.go b/client/internal/peer/notifier.go index 1ee1d32ea..564098bd4 100644 --- a/client/internal/peer/notifier.go +++ b/client/internal/peer/notifier.go @@ -12,6 +12,8 @@ type notifier struct { serverStateLock sync.Mutex listenersLock sync.Mutex listener Listener + peerListWake chan struct{} + peerListStop chan struct{} currentClientState bool lastNotification ClientState lastNumberOfPeers int @@ -62,7 +64,6 @@ func (n *notifier) setNetworkAvailable(available bool) { func (n *notifier) setListener(listener Listener) { n.serverStateLock.Lock() lastNotification := n.effectiveState(n.lastNotification) - numOfPeers := n.lastNumberOfPeers fqdnAddress := n.lastFqdnAddress address := n.lastIPAddress n.serverStateLock.Unlock() @@ -70,17 +71,19 @@ func (n *notifier) setListener(listener Listener) { n.listenersLock.Lock() defer n.listenersLock.Unlock() + n.stopPeerListDelivererLocked() n.listener = listener listener.OnAddressChanged(fqdnAddress, address) notifyListener(listener, lastNotification) - // run on go routine to avoid on Java layer to call go functions on same thread - go listener.OnPeersListChanged(numOfPeers) + n.startPeerListDelivererLocked(listener) + n.wakePeerListDelivererLocked() } func (n *notifier) removeListener() { n.listenersLock.Lock() defer n.listenersLock.Unlock() + n.stopPeerListDelivererLocked() n.listener = nil } @@ -178,15 +181,56 @@ func (n *notifier) peerListChanged(numOfPeers int) { n.serverStateLock.Unlock() n.listenersLock.Lock() - listener := n.listener - n.listenersLock.Unlock() + defer n.listenersLock.Unlock() + n.wakePeerListDelivererLocked() +} - if listener == nil { +func (n *notifier) startPeerListDelivererLocked(listener Listener) { + wake := make(chan struct{}, 1) + stop := make(chan struct{}) + n.peerListWake = wake + n.peerListStop = stop + go n.deliverPeerListChanges(listener, wake, stop) +} + +func (n *notifier) stopPeerListDelivererLocked() { + if n.peerListStop == nil { return } + close(n.peerListStop) + n.peerListStop = nil + n.peerListWake = nil +} - // run on go routine to avoid on Java layer to call go functions on same thread - go listener.OnPeersListChanged(numOfPeers) +func (n *notifier) wakePeerListDelivererLocked() { + if n.peerListWake == nil { + return + } + select { + case n.peerListWake <- struct{}{}: + default: + } +} + +func (n *notifier) deliverPeerListChanges(listener Listener, wake <-chan struct{}, stop <-chan struct{}) { + for { + select { + case <-stop: + return + case <-wake: + } + select { + case <-stop: + return + default: + } + + n.serverStateLock.Lock() + numOfPeers := n.lastNumberOfPeers + n.serverStateLock.Unlock() + + listener.OnPeersListChanged(numOfPeers) + } } func (n *notifier) localAddressChanged(fqdn, address string) { diff --git a/client/internal/peer/notifier_test.go b/client/internal/peer/notifier_test.go index a73016b05..f81866214 100644 --- a/client/internal/peer/notifier_test.go +++ b/client/internal/peer/notifier_test.go @@ -2,7 +2,9 @@ package peer import ( "sync" + "sync/atomic" "testing" + "time" ) type mocListener struct { @@ -115,3 +117,156 @@ func Test_notifier_RemoveListener(t *testing.T) { t.Errorf("invalid state: %d", listener.peers) } } + +type coalescingListener struct { + final int + calls atomic.Int32 + inFlight atomic.Int32 + maxInFlight atomic.Int32 + last atomic.Int32 + done chan struct{} + entered chan struct{} + release chan struct{} + once sync.Once +} + +func (l *coalescingListener) OnStateChanged(ClientState) {} +func (l *coalescingListener) OnConnected() {} +func (l *coalescingListener) OnDisconnected() {} +func (l *coalescingListener) OnConnecting() {} +func (l *coalescingListener) OnDisconnecting() {} +func (l *coalescingListener) OnAddressChanged(string, string) {} + +func (l *coalescingListener) OnPeersListChanged(size int) { + current := l.inFlight.Add(1) + for { + seen := l.maxInFlight.Load() + if current <= seen || l.maxInFlight.CompareAndSwap(seen, current) { + break + } + } + if l.calls.Add(1) == 1 && l.entered != nil { + close(l.entered) + } + if l.release != nil { + <-l.release + } + time.Sleep(time.Millisecond) + l.last.Store(int32(size)) + l.inFlight.Add(-1) + if size == l.final { + l.once.Do(func() { close(l.done) }) + } +} + +func Test_notifier_PeerListChangedCoalesces(t *testing.T) { + const events = 1000 + listener := &coalescingListener{final: events, done: make(chan struct{})} + n := newNotifier() + n.setListener(listener) + + for i := 1; i <= events; i++ { + n.peerListChanged(i) + } + + select { + case <-listener.done: + case <-time.After(5 * time.Second): + t.Fatalf("last peer count not delivered, last seen: %d", listener.last.Load()) + } + + if got := listener.maxInFlight.Load(); got != 1 { + t.Errorf("concurrent deliveries: %d, expected 1", got) + } + if got := listener.calls.Load(); got >= events { + t.Errorf("deliveries not coalesced: %d calls for %d events", got, events) + } +} + +func Test_notifier_SetListenerStopsPreviousDeliverer(t *testing.T) { + old := &coalescingListener{final: -1} + replacement := &coalescingListener{final: 7, done: make(chan struct{})} + n := newNotifier() + n.setListener(old) + oldStop := n.peerListStop + + n.peerListChanged(7) + n.setListener(replacement) + + select { + case <-oldStop: + default: + t.Fatal("old deliverer not stopped on listener replacement") + } + waitFor(t, replacement.done, "replacement listener not notified") +} + +func Test_notifier_RemoveListenerStopsDeliverer(t *testing.T) { + n := newNotifier() + n.setListener(&coalescingListener{final: -1}) + stop := n.peerListStop + + n.removeListener() + + select { + case <-stop: + default: + t.Fatal("deliverer not stopped on listener removal") + } +} + +func Test_notifier_DelivererExitsAfterInFlightCallback(t *testing.T) { + listener := &coalescingListener{ + final: -1, + entered: make(chan struct{}), + release: make(chan struct{}), + } + n := newNotifier() + wake := make(chan struct{}, 1) + stop := make(chan struct{}) + exited := make(chan struct{}) + go func() { + n.deliverPeerListChanges(listener, wake, stop) + close(exited) + }() + + wake <- struct{}{} + waitFor(t, listener.entered, "listener not called") + + n.peerListChanged(7) + wake <- struct{}{} + close(stop) + close(listener.release) + + waitFor(t, exited, "deliverer did not exit after stop") + if got := listener.calls.Load(); got != 1 { + t.Errorf("deliverer ran %d callbacks after stop, expected only the in-flight one", got) + } + if got := listener.last.Load(); got == 7 { + t.Errorf("deliverer delivered the peer count queued after stop") + } +} + +func Test_notifier_DelivererPrefersStopOverPendingWake(t *testing.T) { + listener := &coalescingListener{final: -1} + n := newNotifier() + wake := make(chan struct{}, 1) + stop := make(chan struct{}) + + wake <- struct{}{} + close(stop) + n.deliverPeerListChanges(listener, wake, stop) + + if got := listener.calls.Load(); got != 0 { + t.Errorf("deliverer ran %d callbacks with stop closed, expected 0", got) + } +} + +func waitFor(t *testing.T, ch <-chan struct{}, msg string) { + t.Helper() + select { + case <-ch: + case <-time.After(5 * time.Second): + t.Fatal(msg) + } +} diff --git a/client/internal/peer/status.go b/client/internal/peer/status.go index bf36b944b..6c44178e1 100644 --- a/client/internal/peer/status.go +++ b/client/internal/peer/status.go @@ -18,9 +18,7 @@ import ( "google.golang.org/protobuf/types/known/durationpb" "google.golang.org/protobuf/types/known/timestamppb" - firewall "github.com/netbirdio/netbird/client/firewall/manager" "github.com/netbirdio/netbird/client/iface/configurer" - "github.com/netbirdio/netbird/client/internal/ingressgw" "github.com/netbirdio/netbird/client/internal/relay" "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/route" @@ -161,7 +159,6 @@ type FullStatus struct { RosenpassState RosenpassState Relays []relay.ProbeResult NSGroupStates []NSGroupState - NumOfForwardingRules int LazyConnectionEnabled bool Events []*proto.SystemEvent } @@ -196,6 +193,7 @@ type Status struct { muxRelays sync.RWMutex peers map[string]State ipToKey map[string]string + activeRoutePeers map[route.HAUniqueID]string changeNotify map[string]map[string]*StatusChangeSubscription // map[peerID]map[subscriptionID]*StatusChangeSubscription signalState bool signalError error @@ -246,8 +244,6 @@ type Status struct { // read it without taking mux. networksRevision atomic.Uint64 - ingressGwMgr *ingressgw.Manager - routeIDLookup routeIDLookup wgIface WGIfaceStatus } @@ -257,6 +253,7 @@ func NewRecorder(mgmAddress string) *Status { return &Status{ peers: make(map[string]State), ipToKey: make(map[string]string), + activeRoutePeers: make(map[route.HAUniqueID]string), changeNotify: make(map[string]map[string]*StatusChangeSubscription), eventStreams: make(map[string]chan *proto.SystemEvent), eventQueue: NewEventQueue(eventQueueSize), @@ -274,12 +271,6 @@ func (d *Status) SetRelayMgr(manager *relayClient.Manager) { d.relayMgr = manager } -func (d *Status) SetIngressGwMgr(ingressGwMgr *ingressgw.Manager) { - d.mux.Lock() - defer d.mux.Unlock() - d.ingressGwMgr = ingressGwMgr -} - // ReplaceOfflinePeers replaces func (d *Status) ReplaceOfflinePeers(replacement []State) { d.mux.Lock() @@ -330,18 +321,6 @@ func (d *Status) GetPeer(peerPubKey string) (State, error) { return state, nil } -func (d *Status) PeerByIP(ip string) (string, bool) { - d.mux.RLock() - defer d.mux.RUnlock() - - for _, state := range d.peers { - if state.IP == ip { - return state.FQDN, true - } - } - return "", false -} - // PeerStateByIP returns the full peer State for the given tunnel IP. // Matches against either the IPv4 (State.IP) or IPv6 (State.IPv6) tunnel // address so dual-stack peers are reachable on either family. Only @@ -481,6 +460,24 @@ func (d *Status) RemovePeerStateRoute(peer string, route string) error { return nil } +func (d *Status) AddActiveRoutePeer(haID route.HAUniqueID, peer string) { + d.mux.Lock() + defer d.mux.Unlock() + d.activeRoutePeers[haID] = peer +} + +func (d *Status) RemoveActiveRoutePeer(haID route.HAUniqueID) { + d.mux.Lock() + defer d.mux.Unlock() + delete(d.activeRoutePeers, haID) +} + +func (d *Status) GetActiveRoutePeers() map[route.HAUniqueID]string { + d.mux.RLock() + defer d.mux.RUnlock() + return maps.Clone(d.activeRoutePeers) +} + // CheckRoutes checks if the source and destination addresses are within the same route // and returns the resource ID of the route that contains the addresses func (d *Status) CheckRoutes(ip netip.Addr) ([]byte, bool) { @@ -819,8 +816,8 @@ func (d *Status) SetSessionExpiresAt(deadline time.Time) { // "none" would blank the UI at the exact moment it should say the session // ended. func (d *Status) GetSessionExpiresAt() time.Time { - d.mux.Lock() - defer d.mux.Unlock() + d.mux.RLock() + defer d.mux.RUnlock() return d.sessionExpiresAt } @@ -1143,16 +1140,6 @@ func (d *Status) GetRelayStates() []relay.ProbeResult { return relayStates } -func (d *Status) ForwardingRules() []firewall.ForwardRule { - d.mux.RLock() - defer d.mux.RUnlock() - if d.ingressGwMgr == nil { - return nil - } - - return d.ingressGwMgr.Rules() -} - func (d *Status) GetDNSStates() []NSGroupState { d.mux.RLock() defer d.mux.RUnlock() @@ -1187,7 +1174,6 @@ func (d *Status) GetFullStatus() FullStatus { Relays: d.GetRelayStates(), RosenpassState: d.GetRosenpassState(), NSGroupStates: d.GetDNSStates(), - NumOfForwardingRules: len(d.ForwardingRules()), LazyConnectionEnabled: d.GetLazyConnection(), } @@ -1559,7 +1545,6 @@ func (fs FullStatus) ToProto() *proto.FullStatus { pbFullStatus.LocalPeerState.WgPort = int32(fs.LocalPeerState.WgPort) pbFullStatus.LocalPeerState.RosenpassPermissive = fs.RosenpassState.Permissive pbFullStatus.LocalPeerState.RosenpassEnabled = fs.RosenpassState.Enabled - pbFullStatus.NumberOfForwardingRules = int32(fs.NumOfForwardingRules) pbFullStatus.LazyConnectionEnabled = fs.LazyConnectionEnabled pbFullStatus.LocalPeerState.Networks = maps.Keys(fs.LocalPeerState.Routes) diff --git a/client/internal/peer/status_test.go b/client/internal/peer/status_test.go index 82dff0d6f..b3f01b217 100644 --- a/client/internal/peer/status_test.go +++ b/client/internal/peer/status_test.go @@ -9,6 +9,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/route" ) func TestAddPeer(t *testing.T) { @@ -372,3 +374,24 @@ func TestMarkServerStateDoesNotNotifyWhenUnchanged(t *testing.T) { status.MarkManagementDisconnected(err) assert.False(t, notified(ch), "redundant disconnect should not notify") } + +func TestActiveRoutePeers(t *testing.T) { + status := NewRecorder("https://mgm") + netA := route.HAUniqueID("net-a-10.0.0.0/24") + netB := route.HAUniqueID("net-b-10.0.0.0/24") + + status.AddActiveRoutePeer(netA, "peerA") + status.AddActiveRoutePeer(netB, "peerB") + + active := status.GetActiveRoutePeers() + assert.Equal(t, "peerA", active[netA]) + assert.Equal(t, "peerB", active[netB]) + + status.RemoveActiveRoutePeer(netA) + delete(active, netB) + + active = status.GetActiveRoutePeers() + _, ok := active[netA] + assert.False(t, ok) + assert.Equal(t, "peerB", active[netB]) +} diff --git a/client/internal/peer/worker_ice.go b/client/internal/peer/worker_ice.go index d17f6e693..5979e9bdc 100644 --- a/client/internal/peer/worker_ice.go +++ b/client/internal/peer/worker_ice.go @@ -121,11 +121,8 @@ func (w *WorkerICE) OnNewOffer(remoteOfferAnswer *OfferAnswer) { } } - sessionID, err := NewICESessionID() - if err != nil { - w.log.Errorf("failed to create new session ID: %s", err) - } - w.sessionID = sessionID + // Keep the ID already advertised to the remote. Answers do not get a + // reply, so changing it here makes the next offer restart both sides. w.abandonNegotiation() } @@ -205,6 +202,9 @@ func (w *WorkerICE) Close() { w.muxAgent.Lock() defer w.muxAgent.Unlock() + if w.agent != nil || w.agentConnecting { + w.renewSessionID() + } if w.agent != nil { w.agentDialerCancel() if err := w.agent.Close(); err != nil { @@ -366,16 +366,23 @@ func (w *WorkerICE) closeAgent(agent *icemaker.ThreadSafeAgent, cancel context.C // Only the owner of the current session may reset its state: a stale dial // goroutine waking after a newer attempt must not clobber it. if w.agent == agent { - sessionID, err := NewICESessionID() - if err != nil { - w.log.Errorf("failed to create new session ID: %s", err) - } - w.sessionID = sessionID + w.renewSessionID() w.abandonNegotiation() } return sessionChanged } +// renewSessionID starts a new local session, so the remote treats our next offer +// or answer as a restart. Caller holds muxAgent. +func (w *WorkerICE) renewSessionID() { + sessionID, err := NewICESessionID() + if err != nil { + w.log.Errorf("failed to create new session ID: %s", err) + return + } + w.sessionID = sessionID +} + // abandonNegotiation drops all recorded ICE session state so the worker treats the // next offer as a fresh start instead of a duplicate of a dead negotiation. The // agent and agentConnecting flags must change together: leaving one stale wedges diff --git a/client/internal/peer/worker_ice_session_test.go b/client/internal/peer/worker_ice_session_test.go new file mode 100644 index 000000000..4858e0bc3 --- /dev/null +++ b/client/internal/peer/worker_ice_session_test.go @@ -0,0 +1,375 @@ +package peer + +import ( + "context" + "fmt" + "net" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + icemaker "github.com/netbirdio/netbird/client/internal/peer/ice" +) + +func TestWorkerICE_RemoteRestartPreservesAdvertisedSession(t *testing.T) { + w := newTestWorkerICE(t) + t.Cleanup(w.Close) + w.dialFunc = parkDial + advertised := w.SessionID() + remoteSession := ICESessionID("remote-first") + offer := OfferAnswer{ + IceCredentials: IceCredentials{UFrag: "remoteufrag", Pwd: "remote-password-long-enough"}, + SessionID: &remoteSession, + } + w.OnNewOffer(&offer) + require.True(t, w.InProgress(), "the first remote session must start ICE") + w.muxAgent.Lock() + firstAgent := w.agent + w.muxAgent.Unlock() + + // The same callback handles answers. A changed remote ID must not create + // an unannounced local ID that makes the remote restart on our next offer. + secondSession := ICESessionID("remote-restarted") + answer := offer + answer.SessionID = &secondSession + w.OnNewOffer(&answer) + assert.Equal(t, advertised, w.SessionID(), "following a remote restart must keep our advertised ID") + w.muxAgent.Lock() + secondAgent := w.agent + w.muxAgent.Unlock() + assert.NotSame(t, firstAgent, secondAgent, "the changed remote session must still rebuild ICE") + + w.OnNewOffer(&answer) + w.muxAgent.Lock() + defer w.muxAgent.Unlock() + assert.Same(t, secondAgent, w.agent, "a repeated answer must keep the replacement agent") +} + +func TestWorkerICE_LocalCloseChangesAdvertisedSession(t *testing.T) { + w := newTestWorkerICE(t) + dialStarted := make(chan struct{}) + dialDone := make(chan struct{}) + w.dialFunc = func(ctx context.Context, _ *icemaker.ThreadSafeAgent, _ *OfferAnswer) (net.Conn, error) { + close(dialStarted) + defer close(dialDone) + <-ctx.Done() + return nil, ctx.Err() + } + session := ICESessionID("remote-session") + w.OnNewOffer(&OfferAnswer{ + IceCredentials: IceCredentials{UFrag: "remoteufrag", Pwd: "remote-password-long-enough"}, + SessionID: &session, + }) + <-dialStarted + advertised := w.SessionID() + w.Close() + assert.NotEqual(t, advertised, w.SessionID(), "a local teardown must tell the remote to restart") + closedSession := w.SessionID() + + // The abandoned dial goroutine cleans up after Close returned. + <-dialDone + assert.Never(t, func() bool { return w.SessionID() != closedSession }, 200*time.Millisecond, 10*time.Millisecond, + "the late cleanup of a closed negotiation must not restart again") + w.Close() + assert.Equal(t, closedSession, w.SessionID(), "closing an idle worker must not restart again") +} + +// parkDial stands in for the ICE dial. It never connects and returns once the +// negotiation is abandoned, so a test decides when a negotiation fails. +func parkDial(ctx context.Context, _ *icemaker.ThreadSafeAgent, _ *OfferAnswer) (net.Conn, error) { + <-ctx.Done() + return nil, ctx.Err() +} + +func newTestSessionID(t *testing.T) ICESessionID { + t.Helper() + sid, err := NewICESessionID() + require.NoError(t, err) + return sid +} + +// handshakeSide is one end of a simulated signaling exchange. +type handshakeSide interface { + // message builds the offer or answer the side would send now. + message() OfferAnswer + // receive hands a remote offer or answer to the side's ICE logic. + receive(msg OfferAnswer) + // teardowns counts negotiations the side tore down to follow a remote restart. + teardowns() int + // failAgent ends the side's current negotiation as an ICE failure does. + failAgent() +} + +// workerSide drives a real WorkerICE. +type workerSide struct { + t *testing.T + w *WorkerICE + replaced int +} + +func newWorkerSide(t *testing.T) *workerSide { + t.Helper() + w := newTestWorkerICE(t) + w.dialFunc = parkDial + t.Cleanup(w.Close) + return &workerSide{t: t, w: w} +} + +func (s *workerSide) message() OfferAnswer { + sid := s.w.SessionID() + ufrag, pwd := s.w.GetLocalUserCredentials() + return OfferAnswer{IceCredentials: IceCredentials{UFrag: ufrag, Pwd: pwd}, SessionID: &sid} +} + +func (s *workerSide) receive(msg OfferAnswer) { + before := s.agent() + s.w.OnNewOffer(&msg) + if after := s.agent(); before != nil && after != before { + s.replaced++ + } +} + +func (s *workerSide) teardowns() int { return s.replaced } + +func (s *workerSide) agent() *icemaker.ThreadSafeAgent { + s.w.muxAgent.Lock() + defer s.w.muxAgent.Unlock() + return s.w.agent +} + +// failAgent runs the cleanup the dial goroutine or the Failed state callback +// performs when the current negotiation dies. +func (s *workerSide) failAgent() { + s.t.Helper() + s.w.muxAgent.Lock() + agent, cancel := s.w.agent, s.w.agentDialerCancel + s.w.muxAgent.Unlock() + require.NotNil(s.t, agent, "failing requires a running negotiation") + s.w.closeAgent(agent, cancel) +} + +// legacySide models a remote peer running a release from before this change: +// when it follows a remote restart it also picks a new session ID of its own, +// which it announces only with its next offer or answer. +type legacySide struct { + t *testing.T + sessionID ICESessionID + remoteID ICESessionID + hasAgent bool + replaced int +} + +func newLegacySide(t *testing.T) *legacySide { + return &legacySide{t: t, sessionID: newTestSessionID(t)} +} + +func (s *legacySide) message() OfferAnswer { + sid := s.sessionID + return OfferAnswer{ + IceCredentials: IceCredentials{UFrag: "legacyufrag", Pwd: "legacy-password-long-enough"}, + SessionID: &sid, + } +} + +func (s *legacySide) receive(msg OfferAnswer) { + if msg.SessionID == nil { + s.hasAgent = true + return + } + if s.hasAgent { + if *msg.SessionID == s.remoteID { + return + } + s.replaced++ + s.sessionID = newTestSessionID(s.t) + } + s.hasAgent = true + s.remoteID = *msg.SessionID +} + +func (s *legacySide) teardowns() int { return s.replaced } + +func (s *legacySide) failAgent() { + s.hasAgent = false + s.remoteID = "" + s.sessionID = newTestSessionID(s.t) +} + +// exchange runs one guard-driven round in the order Handshaker.Listen uses: the +// answerer handles the offer and answers with the session ID it holds +// afterwards, and the offerer handles the answer without replying. +func exchange(offerer, answerer handshakeSide) { + answerer.receive(offerer.message()) + offerer.receive(answerer.message()) +} + +// offerPattern decides which side's guard sends the offer in a round. +type offerPattern struct { + name string + picker func(round int, local, remote handshakeSide) (offerer, answerer handshakeSide) +} + +var offerPatterns = []offerPattern{ + { + // A routing peer whose relay is down keeps offering on its own. + name: "local peer offers", + picker: func(_ int, local, remote handshakeSide) (handshakeSide, handshakeSide) { + return local, remote + }, + }, + { + name: "both peers offer", + picker: func(round int, local, remote handshakeSide) (handshakeSide, handshakeSide) { + if round%2 == 0 { + return local, remote + } + return remote, local + }, + }, +} + +// assertSettles runs guard rounds and requires the pair to stop restarting +// each other: at most maxTeardowns in total, and none once half the rounds ran. +func assertSettles(t *testing.T, pattern offerPattern, local, remote handshakeSide, maxTeardowns int) { + t.Helper() + const rounds = 10 + + total := func() int { return local.teardowns() + remote.teardowns() } + start := total() + var halfway int + for round := range rounds { + if round == rounds/2 { + halfway = total() + } + offerer, answerer := pattern.picker(round, local, remote) + exchange(offerer, answerer) + } + + assert.LessOrEqual(t, total()-start, maxTeardowns, "the peers must not keep restarting each other") + assert.Equal(t, halfway, total(), "the negotiation must be stable in the later rounds") +} + +// establish runs the first offer and answer, so both sides negotiate. +func establish(t *testing.T, local, remote handshakeSide) { + t.Helper() + exchange(local, remote) + require.Zero(t, local.teardowns()+remote.teardowns(), "the first exchange must not restart anything") +} + +func TestICESession_SettlesAfterAgentFailure(t *testing.T) { + sides := []struct { + name string + remote func(t *testing.T) handshakeSide + }{ + {name: "current remote", remote: func(t *testing.T) handshakeSide { return newWorkerSide(t) }}, + {name: "legacy remote", remote: func(t *testing.T) handshakeSide { return newLegacySide(t) }}, + } + failures := []struct { + name string + fail func(local, remote handshakeSide) + }{ + {name: "remote agent fails", fail: func(_, remote handshakeSide) { remote.failAgent() }}, + {name: "local agent fails", fail: func(local, _ handshakeSide) { local.failAgent() }}, + {name: "both agents fail", fail: func(local, remote handshakeSide) { + local.failAgent() + remote.failAgent() + }}, + } + + for _, side := range sides { + for _, failure := range failures { + for _, pattern := range offerPatterns { + t.Run(fmt.Sprintf("%s/%s/%s", side.name, failure.name, pattern.name), func(t *testing.T) { + local := newWorkerSide(t) + remote := side.remote(t) + establish(t, local, remote) + + failure.fail(local, remote) + assertSettles(t, pattern, local, remote, 2) + }) + } + } + } +} + +// TestICESession_LocalCloseRestartsRemote covers an explicit teardown, as on a +// WireGuard handshake timeout. The remote must start over as well, or it keeps +// answering from the negotiation this side just abandoned. +func TestICESession_LocalCloseRestartsRemote(t *testing.T) { + for _, pattern := range offerPatterns { + t.Run(pattern.name, func(t *testing.T) { + local := newWorkerSide(t) + remote := newWorkerSide(t) + establish(t, local, remote) + + local.w.Close() + assertSettles(t, pattern, local, remote, 1) + assert.Equal(t, 1, remote.teardowns(), "the remote must restart its negotiation exactly once") + }) + } +} + +func TestICESession_DuplicateMessagesKeepNegotiation(t *testing.T) { + local := newWorkerSide(t) + remote := newWorkerSide(t) + + offer := local.message() + remote.receive(offer) + answer := remote.message() + local.receive(answer) + + // Signaling may deliver the same message again, and a peer answers every + // offer, including repeats of one it already handled. + remote.receive(offer) + local.receive(answer) + local.receive(remote.message()) + + assert.Zero(t, local.teardowns(), "a repeated answer must not restart the negotiation") + assert.Zero(t, remote.teardowns(), "a repeated offer must not restart the negotiation") +} + +// TestICESession_RemoteWithoutSessionIDKeepsNegotiation covers remote peers +// too old to send session IDs: once negotiating, their messages cannot tell a +// restart from a repeat, so they must not tear anything down. +func TestICESession_RemoteWithoutSessionIDKeepsNegotiation(t *testing.T) { + local := newWorkerSide(t) + unversioned := OfferAnswer{IceCredentials: IceCredentials{UFrag: "oldufrag", Pwd: "old-password-long-enough"}} + + local.receive(unversioned) + require.NotNil(t, local.agent(), "a message without a session ID must still start ICE") + advertised := local.w.SessionID() + + for range 3 { + local.receive(unversioned) + } + assert.Zero(t, local.teardowns(), "messages without a session ID must not restart the negotiation") + assert.Equal(t, advertised, local.w.SessionID(), "the advertised session must not change") +} + +// TestWorkerICE_StaleCleanupKeepsAdvertisedSession covers the cleanup of a +// replaced negotiation finishing late, from its dial goroutine or its Closed +// state callback. It must neither pick a new session ID, an unannounced local +// restart, nor disturb the negotiation that replaced it. +func TestWorkerICE_StaleCleanupKeepsAdvertisedSession(t *testing.T) { + local := newWorkerSide(t) + remote := newWorkerSide(t) + establish(t, local, remote) + + local.w.muxAgent.Lock() + oldAgent, oldCancel := local.w.agent, local.w.agentDialerCancel + local.w.muxAgent.Unlock() + + remote.failAgent() + exchange(local, remote) + require.Equal(t, 1, local.teardowns(), "the local side must follow the remote restart") + advertised := local.w.SessionID() + current := local.agent() + + local.w.closeAgent(oldAgent, oldCancel) + + assert.Equal(t, advertised, local.w.SessionID(), "a stale cleanup must not change the advertised session") + assert.Same(t, current, local.agent(), "a stale cleanup must keep the current negotiation") + assertSettles(t, offerPatterns[1], local, remote, 0) +} diff --git a/client/internal/peer/worker_relay.go b/client/internal/peer/worker_relay.go index 0402992c9..694207847 100644 --- a/client/internal/peer/worker_relay.go +++ b/client/internal/peer/worker_relay.go @@ -3,7 +3,6 @@ package peer import ( "context" "errors" - "net" "net/netip" "sync" "sync/atomic" @@ -14,7 +13,7 @@ import ( ) type RelayConnInfo struct { - relayedConn net.Conn + relayedConn *relayClient.Conn rosenpassPubKey []byte rosenpassAddr string } @@ -27,7 +26,7 @@ type WorkerRelay struct { conn *Conn relayManager *relayClient.Manager - relayedConn net.Conn + relayedConn *relayClient.Conn relayLock sync.Mutex relaySupportedOnRemotePeer atomic.Bool @@ -80,12 +79,7 @@ func (w *WorkerRelay) OnNewOffer(remoteOfferAnswer *OfferAnswer) { w.relayedConn = relayedConn w.relayLock.Unlock() - err = w.relayManager.AddCloseListener(srv, w.onRelayClientDisconnected) - if err != nil { - log.Errorf("failed to add close listener: %s", err) - _ = relayedConn.Close() - return - } + go w.watchRelayedConn(relayedConn) w.log.Debugf("peer conn opened via Relay: %s", srv) go w.conn.onRelayConnectionIsReady(RelayConnInfo{ @@ -107,14 +101,21 @@ func (w *WorkerRelay) RelayIsSupportedLocally() bool { return w.relayManager.HasRelayAddress() } +func (w *WorkerRelay) IsTransportConnected() bool { + return w.relayManager.Ready() +} + func (w *WorkerRelay) CloseConn() { w.relayLock.Lock() - defer w.relayLock.Unlock() - if w.relayedConn == nil { + conn := w.relayedConn + w.relayedConn = nil + w.relayLock.Unlock() + + if conn == nil { return } - if err := w.relayedConn.Close(); err != nil { + if err := conn.Close(); err != nil { w.log.Warnf("failed to close relay connection: %v", err) } } @@ -133,6 +134,8 @@ func (w *WorkerRelay) preferredRelayServer(myRelayAddress, remoteRelayAddress st return remoteRelayAddress } -func (w *WorkerRelay) onRelayClientDisconnected() { - go w.conn.onRelayDisconnected() +func (w *WorkerRelay) watchRelayedConn(relayedConn *relayClient.Conn) { + <-relayedConn.Context().Done() + + w.conn.onRelayDisconnected(relayedConn) } diff --git a/client/internal/profilemanager/active_state_test.go b/client/internal/profilemanager/active_state_test.go new file mode 100644 index 000000000..3b7fcd29c --- /dev/null +++ b/client/internal/profilemanager/active_state_test.go @@ -0,0 +1,76 @@ +package profilemanager + +import ( + "fmt" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// Regression test: a concurrent Get and Set of the ActiveProfileState will +// fail on Windows since the write is a temp file renamed over an open file. +// Windows will refuse to replace a file another handle holds open by default. +func TestActiveProfileState_ReadsDoNotBreakAConcurrentWrite(t *testing.T) { + withTempConfigDir(t, func(configDir string) { + withPatchedGlobals(t, configDir, func() { + sm := &ServiceManager{} + require.NoError(t, sm.CreateDefaultProfile()) + require.NoError(t, sm.SetActiveProfileStateToDefault()) + + const switched = ID("0123456789abcdef0123456789abcdef") + const rounds = 50 + + var wg sync.WaitGroup + errs := make(chan error, 128) + + for i := 0; i < 8; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for r := 0; r < rounds; r++ { + state, err := sm.GetActiveProfileState() + if err != nil { + errs <- fmt.Errorf("read: %w", err) + return + } + if state.ID != defaultProfileName && state.ID != switched { + errs <- fmt.Errorf("read: active profile is %q, which no writer wrote", state.ID) + return + } + } + }() + } + + for i := 0; i < 2; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for r := 0; r < rounds; r++ { + id := switched + if r%2 == 0 { + id = defaultProfileName + } + if err := sm.SetActiveProfileState(&ActiveProfileState{ID: id, Username: "testuser"}); err != nil { + errs <- fmt.Errorf("switch: %w", err) + return + } + } + }() + } + + wg.Wait() + close(errs) + + for err := range errs { + assert.NoError(t, err, "a switch and a read of the active profile state must not collide") + } + + state, err := sm.GetActiveProfileState() + require.NoError(t, err) + assert.Contains(t, []ID{defaultProfileName, switched}, state.ID, + "the file holds whichever switch landed last, not a mix of the two") + }) + }) +} diff --git a/client/internal/profilemanager/config.go b/client/internal/profilemanager/config.go index 10c1758d1..ac1b90a62 100644 --- a/client/internal/profilemanager/config.go +++ b/client/internal/profilemanager/config.go @@ -10,7 +10,6 @@ import ( "os" "os/user" "path/filepath" - "reflect" "runtime" "slices" "strings" @@ -58,10 +57,6 @@ var DefaultInterfaceBlacklist = []string{ "Tailscale", "tailscale", "docker", "veth", "br-", "lo", } -// loadMDMPolicy is the package-level indirection used by apply() to read the -// active MDM policy. Tests override this to inject a fake policy. -var loadMDMPolicy = mdm.LoadPolicy - // ConfigInput carries configuration changes to the client type ConfigInput struct { ManagementURL string @@ -202,14 +197,31 @@ type Config struct { MTU uint16 - // policy is the MDM policy that produced the currently-set values for - // any MDM-enforced fields. Set by applyMDMPolicy at the tail of apply() - // and reset on every apply() invocation. Never persisted to disk. - // Callers query enforcement state via Policy() and the mdm.Policy API - // (HasKey, ManagedKeys, IsEmpty). + // probing marks a config that exists only to be compared against and then + // thrown away, so apply() can skip the work that feeds no verdict. + // Unexported, so it never reaches the JSON. + probing bool + + // policy is the MDM policy that produced the currently-set values + // for any MDM-enforced fields. Set by ApplyMDMPolicy on every + // invocation. Never persisted to disk. Callers query enforcement + // state via Policy() and the mdm.Policy API (HasKey, ManagedKeys, + // IsEmpty). policy *mdm.Policy `json:"-"` } +// ApplyMDMPolicy overlays the supplied MDM Policy on top of the current +// Config values and records it as Policy(). The overlay is not reversible: +// an empty Policy only clears the enforcement metadata, so resolve the base +// Config again (from disk or JSON) before applying a changed policy, the way +// the lifecycle owners do on every load. +func (config *Config) ApplyMDMPolicy(policy *mdm.Policy) { + if config == nil { + return + } + config.applyMDMPolicy(policy) +} + // Policy returns the MDM policy applied to this Config. Returns a non-nil // empty Policy when MDM enforcement is inactive; callers can always invoke // HasKey / ManagedKeys / IsEmpty without a nil check. @@ -292,9 +304,11 @@ func fileExists(path string) (bool, error) { return false, err } -// createNewConfig creates a new config generating a new Wireguard key and saving to file -func createNewConfig(input ConfigInput) (*Config, error) { - config := &Config{ +// newConfigSkeleton returns the field values a brand-new profile config starts +// from, before apply() fills in the rest. Shared with the dry-run baseline so +// the two cannot disagree about what "a new config" means. +func newConfigSkeleton() *Config { + return &Config{ // defaults to false only for new (post 0.26) configurations ServerSSHAllowed: util.False(), // Remote jobs are an explicit opt-in and default off, including for @@ -302,6 +316,91 @@ func createNewConfig(input ConfigInput) (*Config, error) { RemoteJobsAllowed: util.False(), WgPort: iface.DefaultWgPort, } +} + +// resolveUnsetDefaults is the single place where an optional field that carries +// no value gets one, and the only place that states what each of those defaults +// is. apply() runs it before it compares anything, and that ordering is the +// point: with the values named, every comparison below it diffs values instead +// of presence. +// +// Presence-based comparison is what broke `netbird up` for a client configured +// through the environment. These fields mean "the effective default" when they +// hold nothing — every consumer already reads a nil as the value resolved here, +// the SSH toggles in engine_ssh.go and the network monitor in +// createEngineConfig — so naming them changes nothing about what runs. But +// while they stayed nil, an input restating the default read as a change, and +// since the CLI sends every flag whose value came from an environment variable +// on each `netbird up`, a client with NB_ENABLE_SSH_ROOT=false restated it +// every time and the update-settings gate refused it. +// +// Filling a field in is not a settings change, so a caller measuring change +// must not read the returned bool as one: see WouldChange, which runs a pass +// for this and discards its verdict. +// +// ServerSSHAllowed is the one field whose default depends on the config's age. +// A brand-new profile gets false from newConfigSkeleton, which runs before +// this, so what is resolved here is only the legacy case: a config written by a +// version that had no such field keeps SSH on, for backwards compatibility. +func (config *Config) resolveUnsetDefaults() (updated bool) { + // Fields that default to false on every platform. + for _, field := range []**bool{ + &config.EnableSSHRoot, + &config.EnableSSHSFTP, + &config.EnableSSHLocalPortForwarding, + &config.EnableSSHRemotePortForwarding, + &config.DisableSSHAuth, + // Remote jobs are an explicit opt-in: unlike SSH, a pre-existing config + // with no value defaults to disabled rather than being turned on. + &config.RemoteJobsAllowed, + } { + if *field == nil { + *field = util.False() + updated = true + } + } + + if config.DisableNotifications == nil { + log.Infof("setting notifications to disabled by default") + config.DisableNotifications = util.True() + updated = true + } + + if config.SSHJWTCacheTTL == nil { + // A zero TTL disables the JWT cache, which is what no value meant. + config.SSHJWTCacheTTL = new(int) + updated = true + } + + if config.NetworkMonitor == nil { + // network monitoring is on by default on windows and darwin clients + enabled := runtime.GOOS == "windows" || runtime.GOOS == "darwin" + config.NetworkMonitor = &enabled + updated = true + } + + if config.ServerSSHAllowed == nil { + if runtime.GOOS == "android" { + // default to disabled SSH on Android for security + log.Infof("setting SSH server to false by default on Android") + config.ServerSSHAllowed = util.False() + } else { + // enables SSH for configs from old versions to preserve backwards compatibility + log.Infof("falling back to enabled SSH server for pre-existing configuration") + config.ServerSSHAllowed = util.True() + } + updated = true + } + + return updated +} + +// createNewConfig resolves a new config in memory, with no identity: whoever +// needs the peer's keys calls EnsureIdentity and persists the result, so a read +// that lands on a missing file cannot hand back a config carrying keys that +// nothing will ever write down. +func createNewConfig(input ConfigInput) (*Config, error) { + config := newConfigSkeleton() if _, err := config.apply(input); err != nil { return nil, err @@ -310,6 +409,52 @@ func createNewConfig(input ConfigInput) (*Config, error) { return config, nil } +// createProvisionedConfig is createNewConfig plus the peer's identity, for the +// callers that go on to persist the config or to connect with it. +func createProvisionedConfig(input ConfigInput) (*Config, error) { + config, err := createNewConfig(input) + if err != nil { + return nil, err + } + + if _, err := config.EnsureIdentity(); err != nil { + return nil, err + } + + return config, nil +} + +// EnsureIdentity generates the keys that identify this peer if the config does +// not carry them yet, reporting whether it had to generate any. +// +// It is deliberately not part of apply(). Everything apply() fills in is a +// default it can recompute on the next read, but a generated key is not: it +// has to be persisted, or the peer comes back with a different WireGuard +// identity and re-registers. Having apply() generate keys is what forced every +// read of a config to write it back — so identity provisioning is its own step +// now, and the callers that perform it write the result out explicitly. +func (config *Config) EnsureIdentity() (bool, error) { + generated := false + + if config.PrivateKey == "" { + log.Infof("generated new Wireguard key") + config.PrivateKey = generateKey() + generated = true + } + + if config.SSHKey == "" { + log.Infof("generated new SSH key") + pem, err := ssh.GeneratePrivateKey(ssh.ED25519) + if err != nil { + return generated, err + } + config.SSHKey = string(pem) + generated = true + } + + return generated, nil +} + func (config *Config) apply(input ConfigInput) (updated bool, err error) { if config.Name != "" { sanitized, err := sanitizeDisplayName(config.Name) @@ -321,6 +466,13 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } } + + // Every optional field gets its value here, before anything below compares + // one. See resolveUnsetDefaults for why that ordering is the point. + if config.resolveUnsetDefaults() { + updated = true + } + if config.ManagementURL == nil { log.Infof("using default Management URL %s", DefaultManagementURL) config.ManagementURL, err = parseURL("Management URL", DefaultManagementURL) @@ -328,20 +480,21 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { return false, err } } - if input.ManagementURL != "" && input.ManagementURL != config.ManagementURL.String() { - log.Infof("new Management URL provided, updated to %#v (old value %#v)", - input.ManagementURL, config.ManagementURL.String()) + // The comparison is on the endpoint the URL addresses, not on its + // spelling: the same endpoint can be written several ways (an implicit + // :443, a trailing slash, a different host case), and treating an + // equivalent URL as new would rewrite the config and report a settings + // change where the configuration does not actually change. + if input.ManagementURL != "" { URL, err := parseURL("Management URL", input.ManagementURL) if err != nil { return false, err } - config.ManagementURL = URL - updated = true - } else if config.ManagementURL == nil { - log.Infof("using default Management URL %s", DefaultManagementURL) - config.ManagementURL, err = parseURL("Management URL", DefaultManagementURL) - if err != nil { - return false, err + if !SameServiceURL(URL, config.ManagementURL) { + log.Infof("new Management URL provided, updated to %#v (old value %#v)", + URL.String(), config.ManagementURL.String()) + config.ManagementURL = URL + updated = true } } @@ -352,31 +505,20 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { return false, err } } - if input.AdminURL != "" && input.AdminURL != config.AdminURL.String() { - log.Infof("new Admin Panel URL provided, updated to %#v (old value %#v)", - input.AdminURL, config.AdminURL.String()) + // The admin panel is opened, not dialed, so unlike the Management URL its + // path is part of what identifies it: a panel served under /netbird is not + // the one served at the root. + if input.AdminURL != "" { newURL, err := parseURL("Admin Panel URL", input.AdminURL) if err != nil { return updated, err } - config.AdminURL = newURL - updated = true - } - - if config.PrivateKey == "" { - log.Infof("generated new Wireguard key") - config.PrivateKey = generateKey() - updated = true - } - - if config.SSHKey == "" { - log.Infof("generated new SSH key") - pem, err := ssh.GeneratePrivateKey(ssh.ED25519) - if err != nil { - return false, err + if !SameServiceURLIncludingPath(newURL, config.AdminURL) { + log.Infof("new Admin Panel URL provided, updated to %#v (old value %#v)", + newURL.String(), config.AdminURL.String()) + config.AdminURL = newURL + updated = true } - config.SSHKey = string(pem) - updated = true } if input.WireguardPort != nil && *input.WireguardPort != config.WgPort { @@ -397,7 +539,14 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.NATExternalIPs != nil && !reflect.DeepEqual(config.NATExternalIPs, input.NATExternalIPs) { + // slices.Equal, not reflect.DeepEqual, and for the same reason the DNS + // labels below use it: DeepEqual calls a nil slice and an empty one + // different, while both mean "no NAT mappings". A profile stores the + // absent list as JSON null and reads it back nil, and `netbird up` sends + // CleanNATExternalIPs — an empty list — whenever NB_EXTERNAL_IP_MAP is set + // to nothing, so the two met on every start and the gate read a no-op as a + // settings change. + if input.NATExternalIPs != nil && !slices.Equal(config.NATExternalIPs, input.NATExternalIPs) { log.Infof("updating NAT External IP [ %s ] (old value: [ %s ])", strings.Join(input.NATExternalIPs, " "), strings.Join(config.NATExternalIPs, " ")) @@ -435,21 +584,12 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.NetworkMonitor != nil && (config.NetworkMonitor == nil || *input.NetworkMonitor != *config.NetworkMonitor) { + if input.NetworkMonitor != nil && *input.NetworkMonitor != *config.NetworkMonitor { log.Infof("switching Network Monitor to %t", *input.NetworkMonitor) config.NetworkMonitor = input.NetworkMonitor updated = true } - if config.NetworkMonitor == nil { - // enable network monitoring by default on windows and darwin clients - if runtime.GOOS == "windows" || runtime.GOOS == "darwin" { - enabled := true - config.NetworkMonitor = &enabled - updated = true - } - } - if input.CustomDNSAddress != nil && string(input.CustomDNSAddress) != config.CustomDNSAddress { log.Infof("updating custom DNS address %#v (old value %#v)", string(input.CustomDNSAddress), config.CustomDNSAddress) @@ -482,7 +622,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.ServerSSHAllowed != nil && (config.ServerSSHAllowed == nil || *input.ServerSSHAllowed != *config.ServerSSHAllowed) { + if input.ServerSSHAllowed != nil && *input.ServerSSHAllowed != *config.ServerSSHAllowed { if *input.ServerSSHAllowed { log.Infof("enabling SSH server") } else { @@ -490,20 +630,9 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { } config.ServerSSHAllowed = input.ServerSSHAllowed updated = true - } else if config.ServerSSHAllowed == nil { - if runtime.GOOS == "android" { - // default to disabled SSH on Android for security - log.Infof("setting SSH server to false by default on Android") - config.ServerSSHAllowed = util.False() - } else { - // enables SSH for configs from old versions to preserve backwards compatibility - log.Infof("falling back to enabled SSH server for pre-existing configuration") - config.ServerSSHAllowed = util.True() - } - updated = true } - if input.RemoteJobsAllowed != nil && (config.RemoteJobsAllowed == nil || *input.RemoteJobsAllowed != *config.RemoteJobsAllowed) { + if input.RemoteJobsAllowed != nil && *input.RemoteJobsAllowed != *config.RemoteJobsAllowed { if *input.RemoteJobsAllowed { log.Infof("enabling remote jobs") } else { @@ -511,14 +640,9 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { } config.RemoteJobsAllowed = input.RemoteJobsAllowed updated = true - } else if config.RemoteJobsAllowed == nil { - // Remote jobs are an explicit opt-in: unlike SSH, a pre-existing config - // with no value defaults to disabled rather than being turned on. - config.RemoteJobsAllowed = util.False() - updated = true } - if input.EnableSSHRoot != nil && (config.EnableSSHRoot == nil || *input.EnableSSHRoot != *config.EnableSSHRoot) { + if input.EnableSSHRoot != nil && *input.EnableSSHRoot != *config.EnableSSHRoot { if *input.EnableSSHRoot { log.Infof("enabling SSH root login") } else { @@ -528,7 +652,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.EnableSSHSFTP != nil && (config.EnableSSHSFTP == nil || *input.EnableSSHSFTP != *config.EnableSSHSFTP) { + if input.EnableSSHSFTP != nil && *input.EnableSSHSFTP != *config.EnableSSHSFTP { if *input.EnableSSHSFTP { log.Infof("enabling SSH SFTP subsystem") } else { @@ -538,7 +662,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.EnableSSHLocalPortForwarding != nil && (config.EnableSSHLocalPortForwarding == nil || *input.EnableSSHLocalPortForwarding != *config.EnableSSHLocalPortForwarding) { + if input.EnableSSHLocalPortForwarding != nil && *input.EnableSSHLocalPortForwarding != *config.EnableSSHLocalPortForwarding { if *input.EnableSSHLocalPortForwarding { log.Infof("enabling SSH local port forwarding") } else { @@ -548,7 +672,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.EnableSSHRemotePortForwarding != nil && (config.EnableSSHRemotePortForwarding == nil || *input.EnableSSHRemotePortForwarding != *config.EnableSSHRemotePortForwarding) { + if input.EnableSSHRemotePortForwarding != nil && *input.EnableSSHRemotePortForwarding != *config.EnableSSHRemotePortForwarding { if *input.EnableSSHRemotePortForwarding { log.Infof("enabling SSH remote port forwarding") } else { @@ -558,7 +682,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.DisableSSHAuth != nil && (config.DisableSSHAuth == nil || *input.DisableSSHAuth != *config.DisableSSHAuth) { + if input.DisableSSHAuth != nil && *input.DisableSSHAuth != *config.DisableSSHAuth { if *input.DisableSSHAuth { log.Infof("disabling SSH authentication") } else { @@ -568,7 +692,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.SSHJWTCacheTTL != nil && (config.SSHJWTCacheTTL == nil || *input.SSHJWTCacheTTL != *config.SSHJWTCacheTTL) { + if input.SSHJWTCacheTTL != nil && *input.SSHJWTCacheTTL != *config.SSHJWTCacheTTL { log.Infof("updating SSH JWT cache TTL to %d seconds", *input.SSHJWTCacheTTL) config.SSHJWTCacheTTL = input.SSHJWTCacheTTL updated = true @@ -651,13 +775,16 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.SyncMessageVersion != nil && *input.SyncMessageVersion != *config.SyncMessageVersion { + // Assigning the pointer, not writing through it: a config that carries no + // version yet would otherwise be a nil dereference, and a panic inside a + // request handler is not a way to fail. + if input.SyncMessageVersion != nil && (config.SyncMessageVersion == nil || *input.SyncMessageVersion != *config.SyncMessageVersion) { log.Infof("setting SyncMessageVersion to %v", *input.SyncMessageVersion) - *config.SyncMessageVersion = *input.SyncMessageVersion + config.SyncMessageVersion = input.SyncMessageVersion updated = true } - if input.DisableNotifications != nil && (config.DisableNotifications == nil || *input.DisableNotifications != *config.DisableNotifications) { + if input.DisableNotifications != nil && *input.DisableNotifications != *config.DisableNotifications { if *input.DisableNotifications { log.Infof("disabling notifications") } else { @@ -667,24 +794,24 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if config.DisableNotifications == nil { - disabled := true - config.DisableNotifications = &disabled - log.Infof("setting notifications to disabled by default") - updated = true - } - - if input.ClientCertKeyPath != "" { + // Compared, not just assigned: restating the path a config already holds + // changes nothing, and reporting it as an update makes a caller that + // re-sends its own configuration look like one asking to change it. + if input.ClientCertKeyPath != "" && input.ClientCertKeyPath != config.ClientCertKeyPath { config.ClientCertKeyPath = input.ClientCertKeyPath updated = true } - if input.ClientCertPath != "" { + if input.ClientCertPath != "" && input.ClientCertPath != config.ClientCertPath { config.ClientCertPath = input.ClientCertPath updated = true } - if config.ClientCertPath != "" && config.ClientCertKeyPath != "" { + // Not on a probe: the loaded pair feeds the connection, never the + // comparison, and this would otherwise run on every gated SetConfig and + // Login — twice per request — including those that are refused or change + // nothing, logging an error per request when the files are missing. + if !config.probing && config.ClientCertPath != "" && config.ClientCertKeyPath != "" { cert, err := tls.LoadX509KeyPair(config.ClientCertPath, config.ClientCertKeyPath) if err != nil { log.Error("Failed to load mTLS cert/key pair: ", err) @@ -712,9 +839,11 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - // MDM is the last override layer: any key present in the policy - // supersedes defaults, on-disk config, env vars and CLI input. - config.applyMDMPolicy(loadMDMPolicy()) + // Initialise the MDM overlay to "no enforcement" so Config.Policy() + // never returns a stale or nil policy on a freshly applied Config. + // Lifecycle owners that want to enforce a real MDM policy invoke + // Config.ApplyMDMPolicy(loader.Load()) after this returns. + config.applyMDMPolicy(mdm.NewPolicy(nil)) return updated, nil } @@ -876,6 +1005,49 @@ func ParseServiceURL(serviceName, serviceURL string) (*url.URL, error) { return parseURL(serviceName, serviceURL) } +// SameServiceURL reports whether two service URLs address the same endpoint: +// same scheme, same host compared case-insensitively as DNS names are, and +// same effective port, where an absent port means the scheme's default. +// +// This is the one comparison every caller deciding "did this URL change?" must +// use. A string comparison answers a different question: "https://host", +// "https://host/" and "https://HOST:443" are one endpoint written three ways, +// and reading them as three values makes a client that restates its own +// management URL look like a client asking to be repointed. A nil operand +// matches only another nil one. +// +// The path plays no part: a management URL is dialed, and only its host and +// port are. util.SameServiceURL is this comparison plus the path, which is +// what SameServiceURLIncludingPath needs and delegates to. +func SameServiceURL(a, b *url.URL) bool { + if a == nil || b == nil { + return a == b + } + + return strings.EqualFold(a.Scheme, b.Scheme) && + strings.EqualFold(a.Hostname(), b.Hostname()) && + util.ServiceURLPort(a) == util.ServiceURLPort(b) +} + +// SameServiceURLIncludingPath is SameServiceURL plus everything a URL carries +// past its endpoint: path, query, fragment and userinfo. +// +// Use it for a URL that gets opened rather than dialed. The admin panel can +// live under a path, so two URLs with the same endpoint and different paths are +// two different panels — where for a URL the client dials over gRPC only the +// endpoint is ever used. Equivalent spellings still compare equal: a missing +// path and "/" are the same root, and so is a trailing slash on any path. +func SameServiceURLIncludingPath(a, b *url.URL) bool { + if a == nil || b == nil { + return a == b + } + + return util.SameServiceURL(a, b) && + a.RawQuery == b.RawQuery && + a.Fragment == b.Fragment && + a.User.String() == b.User.String() +} + func parseURL(serviceName, serviceURL string) (*url.URL, error) { parsedMgmtURL, err := url.ParseRequestURI(serviceURL) if err != nil { @@ -920,6 +1092,84 @@ func isPreSharedKeyHidden(preSharedKey *string) bool { return false } +// WouldChange reports whether applying input would modify any field the +// config persists, leaving the receiver untouched. It is the dry-run half of +// UpdateConfig and reuses the very same diff logic (Config.apply), so a +// caller asking "is this a settings change?" cannot drift from what an +// actual update would do, nor go stale when a new field is added. +// +// A redacted pre-shared key is collapsed to "unset" exactly as +// UpdateOrCreateConfig does, so a UI that round-trips the mask is not read as +// a request for a new key. +// +// A nil receiver means the profile holds no config yet, so the baseline is the +// config the daemon would create for it: input values matching those defaults +// change nothing, anything else does. +func (config *Config) WouldChange(input ConfigInput) (bool, error) { + probe := config.clone() + if probe == nil { + baseline, err := newDryRunBaseline(input.ConfigPath) + if err != nil { + return true, fmt.Errorf("build default config baseline: %w", err) + } + probe = baseline + } + probe.probing = true + + // Normalize before measuring. apply() reports two different things through + // one bool: an input that changed a value, and a field it had to fill in + // because the config carried none. Only the first is a settings change, so + // the filling-in gets a pass of its own whose verdict is discarded, and the + // pass that answers the caller runs against a config with nothing left to + // fill in. + // + // Readers already hand out normalized configs — readConfig applies an empty + // input for this very reason — so this is normally a no-op. But a gate that + // refuses a request must not depend on where its caller got the config + // from, and it must not start reading "this profile predates a field" as + // "the caller asked for a change" the day someone adds one. + if _, err := probe.apply(ConfigInput{ConfigPath: input.ConfigPath}); err != nil { + return true, fmt.Errorf("normalize the config to diff against: %w", err) + } + + if isPreSharedKeyHidden(input.PreSharedKey) { + input.PreSharedKey = nil + } + + return probe.apply(input) +} + +// newDryRunBaseline builds the config a brand-new profile would start from, for +// a dry run to compare an input against. It is createNewConfig without the +// identity: this config exists only to be compared against and thrown away, and +// no ConfigInput field maps to either key. +func newDryRunBaseline(configPath string) (*Config, error) { + baseline := newConfigSkeleton() + + if _, err := baseline.apply(ConfigInput{ConfigPath: configPath}); err != nil { + return nil, err + } + + return baseline, nil +} + +// clone returns a copy of the config that apply can be run against without the +// original observing the writes, or nil for a nil receiver. Only what apply +// mutates in place needs detaching, which is the slices it replaces or appends +// to: every pointer field it touches is reassigned rather than written through, +// and ClientCertKeyPair is only overwritten. +func (config *Config) clone() *Config { + if config == nil { + return nil + } + + probe := *config + probe.IFaceBlackList = slices.Clone(config.IFaceBlackList) + probe.NATExternalIPs = slices.Clone(config.NATExternalIPs) + probe.DNSLabels = slices.Clone(config.DNSLabels) + return &probe +} + // UpdateConfig update existing configuration according to input configuration and return with the configuration func UpdateConfig(input ConfigInput) (*Config, error) { configExists, err := fileExists(input.ConfigPath) @@ -930,6 +1180,14 @@ func UpdateConfig(input ConfigInput) (*Config, error) { return nil, fmt.Errorf("config file %s does not exist", input.ConfigPath) } + // A UI that round-trips the mask GetConfig hands it back is asking to keep + // the stored key, not to set the mask as the new one. UpdateOrCreateConfig + // and DirectUpdateOrCreateConfig already collapse it; this one did not, so + // the same round-trip through SetConfig replaced the key with asterisks. + if isPreSharedKeyHidden(input.PreSharedKey) { + input.PreSharedKey = nil + } + return update(input) } @@ -941,7 +1199,7 @@ func UpdateOrCreateConfig(input ConfigInput) (*Config, error) { } if !configExists { log.Infof("generating new config %s", input.ConfigPath) - cfg, err := createNewConfig(input) + cfg, err := createProvisionedConfig(input) if err != nil { return nil, err } @@ -966,12 +1224,20 @@ func update(input ConfigInput) (*Config, error) { return nil, err } + // A write path is a provisioning point: a stored profile can legitimately + // carry no identity (a mobile logout clears the keys in place), and the + // next config write is what has to mint a new one. Reads leave that alone. + identityGenerated, err := config.EnsureIdentity() + if err != nil { + return nil, err + } + updated, err := config.apply(input) if err != nil { return nil, err } - if updated { + if updated || identityGenerated { if err := util.WriteJson(context.Background(), input.ConfigPath, config); err != nil { return nil, err } @@ -980,8 +1246,8 @@ func update(input ConfigInput) (*Config, error) { return config, nil } -// GetConfig read config file and return with Config and if it was created. Errors out if it does not exist -func GetConfig(configPath string) (*Config, error) { +// GetExistingConfig reads and returns the config if it exists on disk. Fails otherwise. +func GetExistingConfig(configPath string) (*Config, error) { return readConfig(configPath, false) } @@ -1064,17 +1330,27 @@ func UpdateOldManagementURL(ctx context.Context, config *Config, configPath stri return newConfig, nil } -// CreateInMemoryConfig generate a new config but do not write out it to the store +// CreateInMemoryConfig generate a new config but do not write out it to the store. +// It carries an identity: callers connect with what they get back. func CreateInMemoryConfig(input ConfigInput) (*Config, error) { - return createNewConfig(input) + return createProvisionedConfig(input) } -// ReadConfig read config file and return with Config. If it is not exists create a new with default values -func ReadConfig(configPath string) (*Config, error) { +// ReadConfigOrDefault reads the profile config at configPath, or resolves the +// default config in memory when the file does not exist. It never writes, and +// never mints an identity — EnsureIdentity is where that happens, so the +// caller that provisions is also the one that persists. +func ReadConfigOrDefault(configPath string) (*Config, error) { return readConfig(configPath, true) } -// ReadConfig read config file and return with Config. If it is not exists create a new with default values +// readConfig reads the profile config at configPath. createIfMissing resolves a +// default config in memory when the file is absent, rather than erroring. +// +// Reads are pure. This used to write the config back whenever apply() had to +// fill in a default the file was missing, which quietly made every reader a +// writer: a gate deciding whether to refuse a request, a UI listing profiles, +// a mobile getter reading a single preference. func readConfig(configPath string, createIfMissing bool) (*Config, error) { configExists, err := fileExists(configPath) if err != nil { @@ -1092,12 +1368,8 @@ func readConfig(configPath string, createIfMissing bool) (*Config, error) { return nil, err } // initialize through apply() without changes - if changed, err := config.apply(ConfigInput{}); err != nil { + if _, err := config.apply(ConfigInput{}); err != nil { return nil, err - } else if changed { - if err = WriteOutConfig(configPath, config); err != nil { - return nil, err - } } return config, nil @@ -1105,13 +1377,7 @@ func readConfig(configPath string, createIfMissing bool) (*Config, error) { return nil, fmt.Errorf("config file %s does not exist", configPath) } - cfg, err := createNewConfig(ConfigInput{ConfigPath: configPath}) - if err != nil { - return nil, err - } - - err = WriteOutConfig(configPath, cfg) - return cfg, err + return createNewConfig(ConfigInput{ConfigPath: configPath}) } // WriteOutConfig write put the prepared config to the given path @@ -1134,7 +1400,7 @@ func DirectUpdateOrCreateConfig(input ConfigInput) (*Config, error) { } if !configExists { log.Infof("generating new config %s", input.ConfigPath) - cfg, err := createNewConfig(input) + cfg, err := createProvisionedConfig(input) if err != nil { return nil, err } @@ -1161,12 +1427,18 @@ func directUpdate(input ConfigInput) (*Config, error) { return nil, err } + // Same provisioning point as update(); see the note there. + identityGenerated, err := config.EnsureIdentity() + if err != nil { + return nil, err + } + updated, err := config.apply(input) if err != nil { return nil, err } - if updated { + if updated || identityGenerated { if err := util.DirectWriteJson(context.Background(), input.ConfigPath, config); err != nil { return nil, err } @@ -1188,7 +1460,16 @@ func ConfigToJSON(config *Config) (string, error) { // ConfigFromJSON deserializes a JSON string to a Config struct. // This is useful for restoring config from alternative storage mechanisms. -// After unmarshaling, defaults are applied to ensure the config is fully initialized. +// After unmarshaling, defaults are applied to ensure the config is fully +// initialized. +// +// The peer identity is deliberately none of its business, in either direction. +// It does not generate one: a read cannot hand back keys that nothing will +// write down (see ReadConfigOrDefault). Nor does it refuse a document that +// carries none, because a config legitimately has no identity between a logout +// and the next login — mobile logout clears both keys in place — and this is +// also the deserializer the iOS SDK copies a config through. Whoever goes on +// to connect is where an absent identity has to be answered. func ConfigFromJSON(jsonStr string) (*Config, error) { config := &Config{} err := json.Unmarshal([]byte(jsonStr), config) diff --git a/client/internal/profilemanager/config_json_test.go b/client/internal/profilemanager/config_json_test.go new file mode 100644 index 000000000..9a6d820c4 --- /dev/null +++ b/client/internal/profilemanager/config_json_test.go @@ -0,0 +1,44 @@ +package profilemanager + +import ( + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" +) + +// The serialized form is how the tvOS SDK stores a profile and how the iOS SDK +// copies one in memory, so it must round-trip whatever a profile legitimately +// holds — including no identity at all, which is the state mobile logout leaves +// behind when it clears both keys in place. Refusing that document here broke +// logout, profile switching and the login that follows them. +func TestConfigFromJSONRoundTripsALoggedOutProfile(t *testing.T) { + path := filepath.Join(t.TempDir(), "exported.json") + stored, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path, ManagementURL: DefaultManagementURL}) + require.NoError(t, err) + require.NotEmpty(t, stored.PrivateKey, "a provisioned config is the fixture this test starts from") + require.NotEmpty(t, stored.SSHKey) + + exported, err := ConfigToJSON(stored) + require.NoError(t, err) + + restored, err := ConfigFromJSON(exported) + require.NoError(t, err, "a config exported after a login must load") + require.Equal(t, stored.PrivateKey, restored.PrivateKey, "the restored peer is not the stored one") + require.Equal(t, stored.SSHKey, restored.SSHKey) + + // What mobile logout leaves on disk. + loggedOut := stored.clone() + loggedOut.PrivateKey = "" + loggedOut.SSHKey = "" + + document, err := ConfigToJSON(loggedOut) + require.NoError(t, err) + + reloaded, err := ConfigFromJSON(document) + require.NoError(t, err, "a logged-out profile must still load") + require.Empty(t, reloaded.PrivateKey, "loading must not mint a key nothing will write down") + require.Empty(t, reloaded.SSHKey) + require.Equal(t, stored.ManagementURL.String(), reloaded.ManagementURL.String(), + "the rest of the profile survives the logout") +} diff --git a/client/internal/profilemanager/config_mdm.go b/client/internal/profilemanager/config_mdm.go new file mode 100644 index 000000000..25b9f18f7 --- /dev/null +++ b/client/internal/profilemanager/config_mdm.go @@ -0,0 +1,52 @@ +package profilemanager + +import ( + "errors" + "fmt" + + "github.com/netbirdio/netbird/client/mdm" +) + +// ErrMDMManagedFields marks a config change rejected because it diverges from +// MDM-enforced values. +var ErrMDMManagedFields = errors.New("fields managed by MDM cannot be modified") + +// MDMConflicts returns the names of MDM-managed keys whose requested value in +// the ConfigInput differs from the policy-enforced value; a field set to the +// enforced value is a no-op echo, not a conflict. +func MDMConflicts(input ConfigInput, policy *mdm.Policy) []string { + pskGot := input.PreSharedKey + if isPreSharedKeyHidden(pskGot) { + pskGot = nil + } + var port *int64 + if input.WireguardPort != nil { + v := int64(*input.WireguardPort) + port = &v + } + return mdm.ResolveConflicts(policy, []mdm.ConflictCheck{ + mdm.ConflictURL(mdm.KeyManagementURL, input.ManagementURL), + mdm.ConflictStringPtr(mdm.KeyPreSharedKey, pskGot), + mdm.ConflictBool(mdm.KeyRosenpassEnabled, input.RosenpassEnabled), + mdm.ConflictBool(mdm.KeyRosenpassPermissive, input.RosenpassPermissive), + mdm.ConflictBool(mdm.KeyDisableAutoConnect, input.DisableAutoConnect), + mdm.ConflictBool(mdm.KeyAllowServerSSH, input.ServerSSHAllowed), + mdm.ConflictBool(mdm.KeyRemoteJobsAllowed, input.RemoteJobsAllowed), + mdm.ConflictBool(mdm.KeyDisableClientRoutes, input.DisableClientRoutes), + mdm.ConflictBool(mdm.KeyDisableServerRoutes, input.DisableServerRoutes), + mdm.ConflictBool(mdm.KeyBlockInbound, input.BlockInbound), + mdm.ConflictInt64(mdm.KeyWireguardPort, port), + mdm.ConflictBool(mdm.KeyEnableLocalMetrics, input.LocalMetricsEnabled), + mdm.ConflictStringPtr(mdm.KeyLocalMetricsAddress, input.LocalMetricsAddress), + }) +} + +// CheckMDMConflicts returns an ErrMDMManagedFields-wrapped error naming the +// conflicting keys, or nil when the input does not fight the policy. +func CheckMDMConflicts(input ConfigInput, policy *mdm.Policy) error { + conflicts := MDMConflicts(input, policy) + if len(conflicts) == 0 { + return nil + } + return fmt.Errorf("%w: %v", ErrMDMManagedFields, conflicts) +} diff --git a/client/internal/profilemanager/config_mdm_test.go b/client/internal/profilemanager/config_mdm_test.go index f8dfddb33..716b7a553 100644 --- a/client/internal/profilemanager/config_mdm_test.go +++ b/client/internal/profilemanager/config_mdm_test.go @@ -10,24 +10,58 @@ import ( "github.com/netbirdio/netbird/client/mdm" ) -// withMDMPolicy temporarily overrides the package-level loadMDMPolicy hook so -// apply() observes the supplied Policy. The original loader is restored at -// test cleanup. -func withMDMPolicy(t *testing.T, policy *mdm.Policy) { +// fakeFetcher implements mdm.PolicyFetcher returning a pre-set policy +// map. Test helper used to construct a Loader without touching the OS +// or any package-level state. +type fakeFetcher struct{ values map[string]any } + +func (f *fakeFetcher) Fetch() map[string]any { return f.values } + +// loaderFor builds an mdm.Loader whose loadPlatform returns the +// supplied Policy's underlying values. +func loaderFor(policy *mdm.Policy) *mdm.Loader { + if policy == nil || policy.IsEmpty() { + return mdm.NewLoader(&fakeFetcher{values: nil}) + } + values := make(map[string]any) + for _, k := range policy.ManagedKeys() { + if v, ok := policy.GetString(k); ok { + values[k] = v + continue + } + if v, ok := policy.GetInt(k); ok { + values[k] = v + continue + } + if v, ok := policy.GetBool(k); ok { + values[k] = v + continue + } + if v, ok := policy.GetStringSlice(k); ok { + values[k] = v + } + } + return mdm.NewLoader(&fakeFetcher{values: values}) +} + +// configWithMDM is the test convenience that builds a Config via +// UpdateOrCreateConfig and overlays the supplied MDM policy on top — +// mirrors the production pattern (Server.getConfig / Client.applyMDMOverlay) +// where the Loader lives outside Config and the apply step is driven +// by the lifecycle owner. +func configWithMDM(t *testing.T, input ConfigInput, policy *mdm.Policy) *Config { t.Helper() - prev := loadMDMPolicy - loadMDMPolicy = func() *mdm.Policy { return policy } - t.Cleanup(func() { loadMDMPolicy = prev }) + cfg, err := UpdateOrCreateConfig(input) + require.NoError(t, err) + require.NotNil(t, cfg) + cfg.ApplyMDMPolicy(loaderFor(policy).Load()) + return cfg } func TestApply_MDMEmpty_NoEnforcement(t *testing.T) { - withMDMPolicy(t, mdm.NewPolicy(nil)) - - cfg, err := UpdateOrCreateConfig(ConfigInput{ + cfg := configWithMDM(t, ConfigInput{ ConfigPath: filepath.Join(t.TempDir(), "config.json"), - }) - require.NoError(t, err) - require.NotNil(t, cfg) + }, mdm.NewPolicy(nil)) assert.True(t, cfg.Policy().IsEmpty(), "no MDM source ⇒ empty Policy") assert.False(t, cfg.Policy().HasKey(mdm.KeyManagementURL)) @@ -39,18 +73,15 @@ func TestApply_MDMEmpty_NoEnforcement(t *testing.T) { func TestApply_MDMOnly_OverridesDefaults(t *testing.T) { const mdmURL = "https://corp.mdm.example.com:443" - withMDMPolicy(t, mdm.NewPolicy(map[string]any{ + + cfg := configWithMDM(t, ConfigInput{ + ConfigPath: filepath.Join(t.TempDir(), "config.json"), + }, mdm.NewPolicy(map[string]any{ mdm.KeyManagementURL: mdmURL, mdm.KeyDisableClientRoutes: true, mdm.KeyBlockInbound: true, })) - cfg, err := UpdateOrCreateConfig(ConfigInput{ - ConfigPath: filepath.Join(t.TempDir(), "config.json"), - }) - require.NoError(t, err) - require.NotNil(t, cfg) - assert.Equal(t, mdmURL, cfg.ManagementURL.String()) assert.True(t, cfg.DisableClientRoutes) assert.True(t, cfg.BlockInbound) @@ -65,16 +96,12 @@ func TestApply_MDMBeatsCLIInput(t *testing.T) { const mdmURL = "https://mdm.example.com:443" const cliURL = "https://cli.example.com:443" - withMDMPolicy(t, mdm.NewPolicy(map[string]any{ - mdm.KeyManagementURL: mdmURL, - })) - - cfg, err := UpdateOrCreateConfig(ConfigInput{ + cfg := configWithMDM(t, ConfigInput{ ConfigPath: filepath.Join(t.TempDir(), "config.json"), ManagementURL: cliURL, - }) - require.NoError(t, err) - require.NotNil(t, cfg) + }, mdm.NewPolicy(map[string]any{ + mdm.KeyManagementURL: mdmURL, + })) // MDM wins over CLI-supplied management URL. assert.Equal(t, mdmURL, cfg.ManagementURL.String()) @@ -82,16 +109,12 @@ func TestApply_MDMBeatsCLIInput(t *testing.T) { } func TestApply_MDMInvalidURL_KeepsPreviousValue(t *testing.T) { - withMDMPolicy(t, mdm.NewPolicy(map[string]any{ + cfg := configWithMDM(t, ConfigInput{ + ConfigPath: filepath.Join(t.TempDir(), "config.json"), + }, mdm.NewPolicy(map[string]any{ mdm.KeyManagementURL: "not-a-url", })) - cfg, err := UpdateOrCreateConfig(ConfigInput{ - ConfigPath: filepath.Join(t.TempDir(), "config.json"), - }) - require.NoError(t, err) - require.NotNil(t, cfg) - // Invalid MDM URL is logged and skipped: default URL stays in place // to keep the client functional. assert.Equal(t, DefaultManagementURL, cfg.ManagementURL.String()) @@ -106,24 +129,20 @@ func TestApply_MDMBoolKeysOverrideOnDiskValue(t *testing.T) { tmp := filepath.Join(t.TempDir(), "config.json") // Seed without MDM. - withMDMPolicy(t, mdm.NewPolicy(nil)) - _, err := UpdateOrCreateConfig(ConfigInput{ + configWithMDM(t, ConfigInput{ ConfigPath: tmp, DisableClientRoutes: boolPtr(false), RosenpassEnabled: boolPtr(false), - }) - require.NoError(t, err) + }, mdm.NewPolicy(nil)) // Now enable MDM enforcement for these keys. - withMDMPolicy(t, mdm.NewPolicy(map[string]any{ + cfg := configWithMDM(t, ConfigInput{ + ConfigPath: tmp, + }, mdm.NewPolicy(map[string]any{ mdm.KeyDisableClientRoutes: true, mdm.KeyRosenpassEnabled: true, })) - cfg, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: tmp}) - require.NoError(t, err) - require.NotNil(t, cfg) - assert.True(t, cfg.DisableClientRoutes, "MDM override should flip on-disk false to true") assert.True(t, cfg.RosenpassEnabled) assert.True(t, cfg.Policy().HasKey(mdm.KeyDisableClientRoutes)) @@ -134,22 +153,19 @@ func TestApply_MDMLocalMetrics(t *testing.T) { tmp := filepath.Join(t.TempDir(), "config.json") // Seed without MDM. - withMDMPolicy(t, mdm.NewPolicy(nil)) - _, err := UpdateOrCreateConfig(ConfigInput{ + configWithMDM(t, ConfigInput{ ConfigPath: tmp, LocalMetricsEnabled: boolPtr(false), - }) - require.NoError(t, err) + }, mdm.NewPolicy(nil)) - withMDMPolicy(t, mdm.NewPolicy(map[string]any{ + // Now enable MDM enforcement for these keys. + cfg := configWithMDM(t, ConfigInput{ + ConfigPath: tmp, + }, mdm.NewPolicy(map[string]any{ mdm.KeyEnableLocalMetrics: true, mdm.KeyLocalMetricsAddress: "127.0.0.1:9292", })) - cfg, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: tmp}) - require.NoError(t, err) - require.NotNil(t, cfg) - assert.True(t, cfg.LocalMetricsEnabled, "MDM override should flip on-disk false to true") assert.Equal(t, "127.0.0.1:9292", cfg.LocalMetricsAddress) assert.True(t, cfg.Policy().HasKey(mdm.KeyEnableLocalMetrics)) @@ -171,16 +187,12 @@ func TestApply_MDMLazyConnection(t *testing.T) { } for _, c := range cases { t.Run(c.name, func(t *testing.T) { - withMDMPolicy(t, mdm.NewPolicy(map[string]any{ + cfg := configWithMDM(t, ConfigInput{ + ConfigPath: filepath.Join(t.TempDir(), "config.json"), + }, mdm.NewPolicy(map[string]any{ mdm.KeyLazyConnection: c.raw, })) - cfg, err := UpdateOrCreateConfig(ConfigInput{ - ConfigPath: filepath.Join(t.TempDir(), "config.json"), - }) - require.NoError(t, err) - require.NotNil(t, cfg) - assert.Equal(t, c.want, cfg.LazyConnection) assert.True(t, cfg.Policy().HasKey(mdm.KeyLazyConnection)) }) @@ -188,22 +200,83 @@ func TestApply_MDMLazyConnection(t *testing.T) { } func TestApply_MDMPreSharedKeyRedactionSentinelRejected(t *testing.T) { - const maskSentinel = "**********" + const maskSentinel = mdm.PreSharedKeyRedactedSentinel - withMDMPolicy(t, mdm.NewPolicy(map[string]any{ + cfg := configWithMDM(t, ConfigInput{ + ConfigPath: filepath.Join(t.TempDir(), "config.json"), + }, mdm.NewPolicy(map[string]any{ mdm.KeyPreSharedKey: maskSentinel, })) - cfg, err := UpdateOrCreateConfig(ConfigInput{ - ConfigPath: filepath.Join(t.TempDir(), "config.json"), - }) - require.NoError(t, err) - require.NotNil(t, cfg) - // Mask sentinel must not be persisted as the actual PSK. assert.NotEqual(t, maskSentinel, cfg.PreSharedKey) // Key still marked managed so user writes are still rejected. assert.True(t, cfg.Policy().HasKey(mdm.KeyPreSharedKey)) } +func TestMDMConflicts_PreSharedKey(t *testing.T) { + policy := mdm.NewPolicy(map[string]any{ + mdm.KeyPreSharedKey: "mdm-enforced-psk", + }) + empty := "" + sentinel := mdm.PreSharedKeyRedactedSentinel + same := "mdm-enforced-psk" + other := "user-psk" + + tests := []struct { + name string + psk *string + want []string + }{ + {name: "unset", psk: nil, want: nil}, + {name: "explicit empty", psk: &empty, want: []string{mdm.KeyPreSharedKey}}, + {name: "sentinel echo", psk: &sentinel, want: nil}, + {name: "same value", psk: &same, want: nil}, + {name: "divergent", psk: &other, want: []string{mdm.KeyPreSharedKey}}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.want, MDMConflicts(ConfigInput{PreSharedKey: tc.psk}, policy)) + }) + } +} + +func TestMDMConflicts_RemoteJobsAndLocalMetrics(t *testing.T) { + policy := mdm.NewPolicy(map[string]any{ + mdm.KeyRemoteJobsAllowed: false, + mdm.KeyEnableLocalMetrics: true, + mdm.KeyLocalMetricsAddress: "127.0.0.1:9999", + }) + sameAddr := "127.0.0.1:9999" + otherAddr := "0.0.0.0:9999" + emptyAddr := "" + + tests := []struct { + name string + input ConfigInput + want []string + }{ + {name: "unset", input: ConfigInput{}, want: nil}, + {name: "echo", input: ConfigInput{ + RemoteJobsAllowed: boolPtr(false), + LocalMetricsEnabled: boolPtr(true), + LocalMetricsAddress: &sameAddr, + }, want: nil}, + {name: "remote jobs divergent", input: ConfigInput{RemoteJobsAllowed: boolPtr(true)}, want: []string{mdm.KeyRemoteJobsAllowed}}, + {name: "metrics disabled", input: ConfigInput{LocalMetricsEnabled: boolPtr(false)}, want: []string{mdm.KeyEnableLocalMetrics}}, + {name: "metrics address divergent", input: ConfigInput{LocalMetricsAddress: &otherAddr}, want: []string{mdm.KeyLocalMetricsAddress}}, + {name: "metrics address explicit empty", input: ConfigInput{LocalMetricsAddress: &emptyAddr}, want: []string{mdm.KeyLocalMetricsAddress}}, + {name: "all divergent", input: ConfigInput{ + RemoteJobsAllowed: boolPtr(true), + LocalMetricsEnabled: boolPtr(false), + LocalMetricsAddress: &otherAddr, + }, want: []string{mdm.KeyRemoteJobsAllowed, mdm.KeyEnableLocalMetrics, mdm.KeyLocalMetricsAddress}}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.want, MDMConflicts(tc.input, policy)) + }) + } +} + func boolPtr(b bool) *bool { return &b } diff --git a/client/internal/profilemanager/config_optional_fields_test.go b/client/internal/profilemanager/config_optional_fields_test.go new file mode 100644 index 000000000..9b74e2217 --- /dev/null +++ b/client/internal/profilemanager/config_optional_fields_test.go @@ -0,0 +1,131 @@ +package profilemanager + +import ( + "encoding/json" + "os" + "path/filepath" + "reflect" + "testing" + + "github.com/stretchr/testify/require" +) + +// optionalBoolFields lists the *bool fields of Config by name, derived from the +// type so a field added later is covered without touching these tests. +func optionalBoolFields() []string { + pointerToBool := reflect.TypeOf((*bool)(nil)) + + var fields []string + configType := reflect.TypeOf(Config{}) + for i := range configType.NumField() { + field := configType.Field(i) + if field.Type == pointerToBool && field.Tag.Get("json") != "-" { + fields = append(fields, field.Name) + } + } + return fields +} + +func requireNoUnsetOptionalBool(t *testing.T, config *Config, context string) { + t.Helper() + + value := reflect.ValueOf(*config) + for _, name := range optionalBoolFields() { + require.False(t, value.FieldByName(name).IsNil(), + "%s left %s unset, so its readers have to invent a default and a diff of it compares presence instead of value", context, name) + } +} + +// An optional bool must not be tristate. While one can be nil, true or false, +// every reader has to invent the meaning of nil, and — the reason this test +// exists — a diff of the config ends up comparing presence rather than value: +// that is what made the update-settings gate refuse `netbird up` for a client +// restating its own defaults. apply() is where a config becomes complete, so +// the invariant belongs to it: no *bool may come out of apply() unset. +func TestApplyLeavesNoOptionalBoolUnset(t *testing.T) { + require.NotEmpty(t, optionalBoolFields(), "the invariant is only meaningful while Config has optional bools") + + t.Run("a config built from scratch", func(t *testing.T) { + config := newConfigSkeleton() + _, err := config.apply(ConfigInput{}) + require.NoError(t, err) + + requireNoUnsetOptionalBool(t, config, "apply on a new config") + }) + + t.Run("a config file that predates every optional field", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "legacy.json") + require.NoError(t, os.WriteFile(path, []byte(`{"WgIface":"wt0"}`), 0o600)) + + config, err := GetExistingConfig(path) + require.NoError(t, err) + + requireNoUnsetOptionalBool(t, config, "a read of a legacy config") + }) + + t.Run("a config file that stores them as null", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "null.json") + _, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path}) + require.NoError(t, err) + unsetOnDisk(t, path, optionalBoolFields()...) + + config, err := GetExistingConfig(path) + require.NoError(t, err) + + requireNoUnsetOptionalBool(t, config, "a read of a config storing nulls") + }) +} + +// The same invariant on disk: what a write leaves in the file is what the next +// client to read it starts from, so no write may store a null. +func TestNoWriteStoresAnUnsetOptionalBool(t *testing.T) { + requireNoNullOnDisk := func(t *testing.T, path string, context string) { + t.Helper() + + raw, err := os.ReadFile(path) + require.NoError(t, err) + + var stored map[string]json.RawMessage + require.NoError(t, json.Unmarshal(raw, &stored)) + + for _, name := range optionalBoolFields() { + value, present := stored[name] + require.True(t, present, "%s did not store %s at all", context, name) + require.NotEqual(t, "null", string(value), "%s stored %s as null", context, name) + } + } + + t.Run("UpdateOrCreateConfig", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "created.json") + _, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path, ManagementURL: DefaultManagementURL}) + require.NoError(t, err) + + requireNoNullOnDisk(t, path, "UpdateOrCreateConfig") + }) + + t.Run("UpdateConfig over a config storing nulls", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "stored.json") + _, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path}) + require.NoError(t, err) + unsetOnDisk(t, path, optionalBoolFields()...) + + _, err = UpdateConfig(ConfigInput{ConfigPath: path, ManagementURL: "https://mgmt.example.com"}) + require.NoError(t, err) + + requireNoNullOnDisk(t, path, "UpdateConfig") + }) + + // Renaming used to copy the file back through a bare Unmarshal, which + // preserved the nulls a pre-fix client had written. + t.Run("RenameProfile", func(t *testing.T) { + withTestSM(t, func(sm *ServiceManager, username string) { + created, err := sm.AddProfile("work", username) + require.NoError(t, err) + unsetOnDisk(t, created.Path, optionalBoolFields()...) + + require.NoError(t, sm.RenameProfile(created.ID, username, "office")) + + requireNoNullOnDisk(t, created.Path, "RenameProfile") + }) + }) +} diff --git a/client/internal/profilemanager/config_probe_test.go b/client/internal/profilemanager/config_probe_test.go new file mode 100644 index 000000000..35a179a84 --- /dev/null +++ b/client/internal/profilemanager/config_probe_test.go @@ -0,0 +1,96 @@ +package profilemanager + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "os" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// writeCertPair writes a throwaway certificate and key, so apply() has +// something real to load rather than a missing file it would only log about. +func writeCertPair(t *testing.T) (certPath, keyPath string) { + t.Helper() + + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + + template := x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "probe-test"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + } + der, err := x509.CreateCertificate(rand.Reader, &template, &template, &key.PublicKey, key) + require.NoError(t, err) + + keyDER, err := x509.MarshalECPrivateKey(key) + require.NoError(t, err) + + dir := t.TempDir() + certPath = filepath.Join(dir, "client.crt") + keyPath = filepath.Join(dir, "client.key") + require.NoError(t, os.WriteFile(certPath, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}), 0o600)) + require.NoError(t, os.WriteFile(keyPath, pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}), 0o600)) + return certPath, keyPath +} + +// The dry run behind the update-settings gate must not read the mTLS pair off +// disk. The loaded pair feeds the connection, never the comparison, and the +// gate runs it on every SetConfig and Login — twice per request — including the +// ones it refuses. +func TestProbeDoesNotLoadTheCertificatePair(t *testing.T) { + certPath, keyPath := writeCertPair(t) + + t.Run("a real apply loads it", func(t *testing.T) { + config := newConfigSkeleton() + config.ClientCertPath, config.ClientCertKeyPath = certPath, keyPath + + _, err := config.apply(ConfigInput{}) + require.NoError(t, err) + require.NotNil(t, config.ClientCertKeyPair, "the connection would have no client certificate") + }) + + t.Run("a probe does not", func(t *testing.T) { + config := newConfigSkeleton() + config.ClientCertPath, config.ClientCertKeyPath = certPath, keyPath + config.probing = true + + _, err := config.apply(ConfigInput{}) + require.NoError(t, err) + require.Nil(t, config.ClientCertKeyPair, "the dry run read the certificate off disk") + }) + + // And the verdict is the same either way, which is the only thing the gate + // asks of the probe. + t.Run("the verdict is unaffected", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "mtls.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: DefaultManagementURL, + ClientCertPath: certPath, + ClientCertKeyPath: keyPath, + }) + require.NoError(t, err) + + stored, err := GetExistingConfig(path) + require.NoError(t, err) + + changed, err := stored.WouldChange(ConfigInput{ClientCertPath: certPath, ClientCertKeyPath: keyPath}) + require.NoError(t, err) + require.False(t, changed, "restating the stored certificate paths is not a change") + + changed, err = stored.WouldChange(ConfigInput{ClientCertPath: filepath.Join(t.TempDir(), "other.crt")}) + require.NoError(t, err) + require.True(t, changed, "a different certificate path is a change") + }) +} diff --git a/client/internal/profilemanager/config_test.go b/client/internal/profilemanager/config_test.go index 248920b5e..a461aa71f 100644 --- a/client/internal/profilemanager/config_test.go +++ b/client/internal/profilemanager/config_test.go @@ -196,7 +196,7 @@ func TestWireguardPortZeroExplicit(t *testing.T) { assert.Equal(t, 0, config.WgPort, "WgPort should be 0 when explicitly set by user") // Verify it persists - readConfig, err := GetConfig(configPath) + readConfig, err := GetExistingConfig(configPath) require.NoError(t, err) assert.Equal(t, 0, readConfig.WgPort, "WgPort should remain 0 after reading from file") } diff --git a/client/internal/profilemanager/config_would_change_test.go b/client/internal/profilemanager/config_would_change_test.go new file mode 100644 index 000000000..6b140030f --- /dev/null +++ b/client/internal/profilemanager/config_would_change_test.go @@ -0,0 +1,529 @@ +package profilemanager + +import ( + "encoding/json" + "os" + "path/filepath" + "runtime" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/iface" + "github.com/netbirdio/netbird/shared/management/domain" +) + +func seededConfig(t *testing.T) *Config { + t.Helper() + + path := filepath.Join(t.TempDir(), "seeded.json") + cfg, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: "https://api.netbird.io:443", + PreSharedKey: strPointer("stored-key"), + }) + require.NoError(t, err) + return cfg +} + +func strPointer(s string) *string { return &s } + +func intPtr(i int) *int { return &i } + +func TestWouldChange(t *testing.T) { + tests := []struct { + name string + input ConfigInput + want bool + }{ + {name: "empty input", input: ConfigInput{}, want: false}, + {name: "same management URL", input: ConfigInput{ManagementURL: "https://api.netbird.io:443"}, want: false}, + {name: "management URL without its default port", input: ConfigInput{ManagementURL: "https://api.netbird.io"}, want: false}, + {name: "different management URL", input: ConfigInput{ManagementURL: "https://other.example:443"}, want: true}, + {name: "same pre-shared key", input: ConfigInput{PreSharedKey: strPointer("stored-key")}, want: false}, + {name: "redacted pre-shared key", input: ConfigInput{PreSharedKey: strPointer("**********")}, want: false}, + {name: "different pre-shared key", input: ConfigInput{PreSharedKey: strPointer("other-key")}, want: true}, + {name: "new interface blacklist entry", input: ConfigInput{ExtraIFaceBlackList: []string{"nb-probe0"}}, want: true}, + {name: "blacklist entry already present", input: ConfigInput{ExtraIFaceBlackList: []string{"lo"}}, want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := seededConfig(t) + + changed, err := cfg.WouldChange(tt.input) + require.NoError(t, err) + require.Equal(t, tt.want, changed) + }) + } +} + +// The dry run must not be observable on the config it is run against: it +// decides whether a write is allowed, it does not perform one. +func TestWouldChangeLeavesTheConfigAlone(t *testing.T) { + cfg := seededConfig(t) + blacklist := len(cfg.IFaceBlackList) + + changed, err := cfg.WouldChange(ConfigInput{ + ManagementURL: "https://other.example:443", + PreSharedKey: strPointer("other-key"), + ExtraIFaceBlackList: []string{"nb-probe0"}, + DNSLabels: domain.FromPunycodeList([]string{"probe"}), + NATExternalIPs: []string{"1.2.3.4"}, + }) + require.NoError(t, err) + require.True(t, changed) + + require.Equal(t, "https://api.netbird.io:443", cfg.ManagementURL.String()) + require.Equal(t, "stored-key", cfg.PreSharedKey) + require.Len(t, cfg.IFaceBlackList, blacklist) + require.Empty(t, cfg.DNSLabels) + require.Empty(t, cfg.NATExternalIPs) +} + +// A nil config means the profile holds nothing yet, so the baseline is what +// the daemon would create for it. +func TestWouldChangeWithoutAStoredConfig(t *testing.T) { + var cfg *Config + + changed, err := cfg.WouldChange(ConfigInput{}) + require.NoError(t, err) + require.False(t, changed, "a request carrying nothing cannot change anything") + + changed, err = cfg.WouldChange(ConfigInput{ManagementURL: DefaultManagementURL}) + require.NoError(t, err) + require.False(t, changed, "the default management URL is what would be written anyway") + + changed, err = cfg.WouldChange(ConfigInput{ManagementURL: "https://other.example:443"}) + require.NoError(t, err) + require.True(t, changed) +} + +func TestWouldChangeReportsAnInvalidInput(t *testing.T) { + cfg := seededConfig(t) + + _, err := cfg.WouldChange(ConfigInput{ManagementURL: "not-a-url"}) + require.Error(t, err) +} + +// Reads must not write. A config file missing a field apply() fills in (MTU, +// here) is what used to trigger the write-back. +func TestReadsDoNotWriteTheConfigBack(t *testing.T) { + denormalized := []byte(`{"WgIface":"wt0"}`) + + for name, read := range map[string]func(string) (*Config, error){ + "GetExistingConfig": GetExistingConfig, + "ReadConfigOrDefault": ReadConfigOrDefault, + } { + t.Run(name, func(t *testing.T) { + path := filepath.Join(t.TempDir(), "profile.json") + require.NoError(t, os.WriteFile(path, denormalized, 0o600)) + + cfg, err := read(path) + require.NoError(t, err) + require.Equal(t, uint16(iface.DefaultMTU), cfg.MTU, "the returned config is still normalized in memory") + require.Empty(t, cfg.PrivateKey, "a read must not mint an identity either") + + after, err := os.ReadFile(path) + require.NoError(t, err) + require.Equal(t, string(denormalized), string(after), "%s rewrote the config file", name) + }) + } +} + +// ReadConfigOrDefault resolves a default config for a profile that has no file +// yet, and that must not create the file either. +func TestReadConfigDoesNotCreateTheFile(t *testing.T) { + path := filepath.Join(t.TempDir(), "absent.json") + + cfg, err := ReadConfigOrDefault(path) + require.NoError(t, err) + require.Equal(t, DefaultManagementURL, cfg.ManagementURL.String()) + + _, err = os.Stat(path) + require.True(t, os.IsNotExist(err), "ReadConfigOrDefault created the config file") +} + +// The identity is the one thing a read cannot recompute, so it is provisioned +// on request and its caller persists it. +func TestEnsureIdentity(t *testing.T) { + cfg := newConfigSkeleton() + + generated, err := cfg.EnsureIdentity() + require.NoError(t, err) + require.True(t, generated) + require.NotEmpty(t, cfg.PrivateKey) + require.NotEmpty(t, cfg.SSHKey) + + key := cfg.PrivateKey + generated, err = cfg.EnsureIdentity() + require.NoError(t, err) + require.False(t, generated, "a config that already has an identity keeps it") + require.Equal(t, key, cfg.PrivateKey) +} + +// One endpoint written several ways is one endpoint. A gate that compared +// spellings refused a client restating its own management URL with a trailing +// slash, which is a normal way to write it. +func TestSameServiceURL(t *testing.T) { + tests := []struct { + a, b string + want bool + }{ + {a: "https://mgmt.example.com", b: "https://mgmt.example.com:443", want: true}, + {a: "https://mgmt.example.com", b: "https://mgmt.example.com/", want: true}, + {a: "https://mgmt.example.com/", b: "https://mgmt.example.com:443/", want: true}, + {a: "https://MGMT.example.com", b: "https://mgmt.example.com", want: true}, + {a: "http://mgmt.example.com", b: "http://mgmt.example.com:80", want: true}, + {a: "https://mgmt.example.com", b: "http://mgmt.example.com", want: false}, + {a: "https://mgmt.example.com", b: "https://mgmt.example.com:8443", want: false}, + {a: "https://mgmt.example.com", b: "https://other.example.com", want: false}, + } + + for _, tt := range tests { + t.Run(tt.a+" vs "+tt.b, func(t *testing.T) { + a, err := ParseServiceURL("a", tt.a) + require.NoError(t, err) + b, err := ParseServiceURL("b", tt.b) + require.NoError(t, err) + + require.Equal(t, tt.want, SameServiceURL(a, b)) + require.Equal(t, tt.want, SameServiceURL(b, a), "the comparison must be symmetric") + }) + } +} + +// The same spellings, through the dry run the update-settings gate uses. +func TestWouldChangeIgnoresURLSpelling(t *testing.T) { + path := filepath.Join(t.TempDir(), "seeded.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: "https://mgmt.example.com", + }) + require.NoError(t, err) + + cfg, err := GetExistingConfig(path) + require.NoError(t, err) + + for _, spelling := range []string{ + "https://mgmt.example.com", + "https://mgmt.example.com/", + "https://mgmt.example.com:443", + "https://mgmt.example.com:443/", + "https://MGMT.example.com", + } { + changed, err := cfg.WouldChange(ConfigInput{ManagementURL: spelling}) + require.NoError(t, err) + require.False(t, changed, "%q is the stored endpoint written differently", spelling) + } + + changed, err := cfg.WouldChange(ConfigInput{ManagementURL: "https://mgmt.example.com:8443"}) + require.NoError(t, err) + require.True(t, changed, "a different port is a different endpoint") +} + +// The dry-run baseline exists to be compared against and discarded, so it must +// not mint keys — the CLI's login backoff loop would otherwise log a fresh +// "generated new Wireguard key" on every attempt. +func TestDryRunBaselineDoesNotGenerateKeys(t *testing.T) { + baseline, err := newDryRunBaseline(filepath.Join(t.TempDir(), "absent.json")) + require.NoError(t, err) + + require.Empty(t, baseline.PrivateKey, "generated a WireGuard key for a throwaway config") + require.Empty(t, baseline.SSHKey, "generated an SSH key for a throwaway config") + + // Everything the comparison actually looks at is still the default config. + require.Equal(t, DefaultManagementURL, baseline.ManagementURL.String()) + require.Equal(t, uint16(iface.DefaultMTU), baseline.MTU) + require.Equal(t, iface.DefaultWgPort, baseline.WgPort) +} + +// A stored profile can carry no identity — a mobile logout clears the keys in +// place — so the next config write has to mint one, which is what keeps the +// following login from dialing management with an empty key. +func TestUpdateConfigProvisionsAMissingIdentity(t *testing.T) { + path := filepath.Join(t.TempDir(), "logged-out.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: "https://api.netbird.io:443", + }) + require.NoError(t, err) + + // Stand in for the logout, which zeroes the keys and writes the config out. + loggedOut, err := GetExistingConfig(path) + require.NoError(t, err) + loggedOut.PrivateKey = "" + loggedOut.SSHKey = "" + require.NoError(t, WriteOutConfig(path, loggedOut)) + + cfg, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path}) + require.NoError(t, err) + require.NotEmpty(t, cfg.PrivateKey, "the write path did not provision an identity") + require.NotEmpty(t, cfg.SSHKey) + + persisted, err := GetExistingConfig(path) + require.NoError(t, err) + require.Equal(t, cfg.PrivateKey, persisted.PrivateKey, "the provisioned identity was not persisted") +} + +// A config that carries no sync message version must not make the dry run +// panic: the gate runs inside a request handler, where failing closed is the +// worst acceptable outcome. +func TestWouldChangeWithoutAStoredSyncMessageVersion(t *testing.T) { + cfg := seededConfig(t) + require.Nil(t, cfg.SyncMessageVersion, "the fixture is only useful while the field starts out unset") + + version := 2 + changed, err := cfg.WouldChange(ConfigInput{SyncMessageVersion: &version}) + require.NoError(t, err) + require.True(t, changed) + require.Nil(t, cfg.SyncMessageVersion, "the dry run set the version on the stored config") +} + +// Restating the certificate paths a config already holds is not a change, for +// the same reason restating any other value is not. +func TestWouldChangeIgnoresRestatedCertificatePaths(t *testing.T) { + path := filepath.Join(t.TempDir(), "mtls.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: "https://api.netbird.io:443", + ClientCertPath: "/etc/netbird/client.crt", + ClientCertKeyPath: "/etc/netbird/client.key", + }) + require.NoError(t, err) + + cfg, err := GetExistingConfig(path) + require.NoError(t, err) + + changed, err := cfg.WouldChange(ConfigInput{ + ClientCertPath: "/etc/netbird/client.crt", + ClientCertKeyPath: "/etc/netbird/client.key", + }) + require.NoError(t, err) + require.False(t, changed, "the stored certificate paths were restated") + + changed, err = cfg.WouldChange(ConfigInput{ClientCertPath: "/etc/netbird/other.crt"}) + require.NoError(t, err) + require.True(t, changed, "a different certificate path is a change") +} + +// A read that lands on a missing file must not hand back keys: nothing would +// write them down, so the caller would connect with an identity that changes on +// the next run and registers a second peer. +func TestReadConfigOrDefaultCarriesNoIdentity(t *testing.T) { + cfg, err := ReadConfigOrDefault(filepath.Join(t.TempDir(), "absent.json")) + require.NoError(t, err) + + require.Empty(t, cfg.PrivateKey, "a read minted a WireGuard key") + require.Empty(t, cfg.SSHKey, "a read minted an SSH key") + + // So the caller's own EnsureIdentity is the one that reports the work, and + // therefore the one that triggers the write. + generated, err := cfg.EnsureIdentity() + require.NoError(t, err) + require.True(t, generated, "the provisioning caller could not tell it had to persist the identity") +} + +// CreateInMemoryConfig is the opposite contract: its callers connect with what +// they get back, so it does carry an identity. +func TestCreateInMemoryConfigCarriesAnIdentity(t *testing.T) { + cfg, err := CreateInMemoryConfig(ConfigInput{ManagementURL: "https://api.netbird.io:443"}) + require.NoError(t, err) + + require.NotEmpty(t, cfg.PrivateKey) + require.NotEmpty(t, cfg.SSHKey) +} + +// The admin panel is opened, not dialed, so its path identifies it. Comparing +// it as a bare endpoint left a custom panel URL unable to change. +func TestAdminURLPathIsPartOfTheIdentity(t *testing.T) { + path := filepath.Join(t.TempDir(), "panel.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + AdminURL: "https://app.example.com/netbird", + }) + require.NoError(t, err) + + cfg, err := GetExistingConfig(path) + require.NoError(t, err) + require.Equal(t, "https://app.example.com:443/netbird", cfg.AdminURL.String()) + + // Equivalent spellings of the same panel are still not a change. + for _, same := range []string{ + "https://app.example.com/netbird", + "https://app.example.com:443/netbird", + "https://app.example.com/netbird/", + "https://APP.example.com/netbird", + } { + changed, err := cfg.WouldChange(ConfigInput{AdminURL: same}) + require.NoError(t, err) + require.False(t, changed, "%q is the stored panel written differently", same) + } + + // A different path is a different panel, and it must be persisted. + changed, err := cfg.WouldChange(ConfigInput{AdminURL: "https://app.example.com/other"}) + require.NoError(t, err) + require.True(t, changed, "a different panel path is a change") + + updated, err := UpdateConfig(ConfigInput{ConfigPath: path, AdminURL: "https://app.example.com/other"}) + require.NoError(t, err) + require.Equal(t, "https://app.example.com:443/other", updated.AdminURL.String(), "the new panel path was not persisted") +} + +// unsetOnDisk rewrites the stored config so the named fields carry a JSON null, +// which is how a profile written before apply() resolved them looks on disk. +// It synthesizes that state: no write produces it any more. +func unsetOnDisk(t *testing.T, path string, fields ...string) { + t.Helper() + + raw, err := os.ReadFile(path) + require.NoError(t, err) + + var stored map[string]json.RawMessage + require.NoError(t, json.Unmarshal(raw, &stored)) + + for _, field := range fields { + _, present := stored[field] + require.True(t, present, "%s is not a field of the stored config", field) + stored[field] = json.RawMessage("null") + } + + rewritten, err := json.Marshal(stored) + require.NoError(t, err) + require.NoError(t, os.WriteFile(path, rewritten, 0600)) +} + +// Seven fields mean "the effective default" when they hold no value, and every +// profile written before apply() resolved them holds them as null. Restating +// that default is asking for no change — and the CLI restates it on every +// `netbird up`, because a flag set through an environment variable is a flag +// pflag reports as Changed. Judging those restatements as changes made the +// update-settings gate refuse `netbird up` outright for a client configured +// through the environment, which is the shape of a Kubernetes deployment. +// +// A login now writes those fields set, so the fixture puts the null state back +// on disk with unsetOnDisk instead of getting it from a login. +func TestWouldChangeIgnoresRestatedDefaultsOfUnsetFields(t *testing.T) { + networkMonitorDefault := runtime.GOOS == "windows" || runtime.GOOS == "darwin" + + tests := []struct { + field string + theDefault ConfigInput + theOtherWay ConfigInput + }{ + {"EnableSSHRoot", + ConfigInput{EnableSSHRoot: boolPtr(false)}, ConfigInput{EnableSSHRoot: boolPtr(true)}}, + {"EnableSSHSFTP", + ConfigInput{EnableSSHSFTP: boolPtr(false)}, ConfigInput{EnableSSHSFTP: boolPtr(true)}}, + {"EnableSSHLocalPortForwarding", + ConfigInput{EnableSSHLocalPortForwarding: boolPtr(false)}, ConfigInput{EnableSSHLocalPortForwarding: boolPtr(true)}}, + {"EnableSSHRemotePortForwarding", + ConfigInput{EnableSSHRemotePortForwarding: boolPtr(false)}, ConfigInput{EnableSSHRemotePortForwarding: boolPtr(true)}}, + {"DisableSSHAuth", + ConfigInput{DisableSSHAuth: boolPtr(false)}, ConfigInput{DisableSSHAuth: boolPtr(true)}}, + {"SSHJWTCacheTTL", + ConfigInput{SSHJWTCacheTTL: intPtr(0)}, ConfigInput{SSHJWTCacheTTL: intPtr(300)}}, + {"NetworkMonitor", + ConfigInput{NetworkMonitor: boolPtr(networkMonitorDefault)}, ConfigInput{NetworkMonitor: boolPtr(!networkMonitorDefault)}}, + } + + for _, tt := range tests { + t.Run(tt.field, func(t *testing.T) { + path := filepath.Join(t.TempDir(), "unset.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: "https://api.netbird.io:443", + }) + require.NoError(t, err) + unsetOnDisk(t, path, tt.field) + + cfg, err := GetExistingConfig(path) + require.NoError(t, err) + + changed, err := cfg.WouldChange(tt.theDefault) + require.NoError(t, err) + require.False(t, changed, "restating the default of an unset %s was judged a change", tt.field) + + // The gate still has to refuse a request that does ask for something. + changed, err = cfg.WouldChange(tt.theOtherWay) + require.NoError(t, err) + require.True(t, changed, "asking for a non-default %s is a change", tt.field) + }) + } +} + +// The verdict must not depend on where the caller got the config from. Readers +// normalize what they hand out, but apply() signals "I filled in a default" +// through the same bool as "the input changed something", so a config that +// never passed through a read would otherwise report a change for an input +// that asks for nothing. +func TestWouldChangeNormalizesBeforeMeasuring(t *testing.T) { + rawConfig := func(t *testing.T) *Config { + t.Helper() + + cfg := &Config{WgIface: iface.WgInterfaceDefault} + require.Nil(t, cfg.ServerSSHAllowed, "the fixture is only useful while the config is not normalized") + require.Nil(t, cfg.EnableSSHRoot) + require.Empty(t, cfg.IFaceBlackList) + return cfg + } + + changed, err := rawConfig(t).WouldChange(ConfigInput{}) + require.NoError(t, err) + require.False(t, changed, "an input carrying nothing cannot change anything") + + changed, err = rawConfig(t).WouldChange(ConfigInput{EnableSSHRoot: boolPtr(false)}) + require.NoError(t, err) + require.False(t, changed, "the default of a field the config never held is not a change") + + changed, err = rawConfig(t).WouldChange(ConfigInput{EnableSSHRoot: boolPtr(true)}) + require.NoError(t, err) + require.True(t, changed, "a non-default value is still a change") +} + +// A zero-padded port addresses the same port. The normalization itself belongs +// to util.ServiceURLPort and is tested there; this asserts that the comparison +// this package hands its callers inherits it. +func TestServiceURLPortIsNormalizedNumerically(t *testing.T) { + padded, err := ParseServiceURL("padded", "https://mgmt.example.com:0443") + require.NoError(t, err) + plain, err := ParseServiceURL("plain", "https://mgmt.example.com:443") + require.NoError(t, err) + + require.True(t, SameServiceURL(padded, plain)) +} + +// A list the profile does not have and a list the request empties are the same +// thing: no NAT mappings, no DNS labels. The profile stores an absent list as +// JSON null and reads it back as a nil slice, while `netbird up` sends the +// emptied list — CleanNATExternalIPs / CleanDNSLabels — whenever the matching +// environment variable is set to nothing, which a deployment template does by +// default. Judging nil and empty as different made the gate refuse that start, +// which is the very deadlock this branch exists to remove, on another field. +func TestWouldChangeIgnoresAnEmptiedListThatWasAlreadyAbsent(t *testing.T) { + path := filepath.Join(t.TempDir(), "lists.json") + _, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path, ManagementURL: DefaultManagementURL}) + require.NoError(t, err) + + stored, err := GetExistingConfig(path) + require.NoError(t, err) + require.Nil(t, stored.NATExternalIPs, "the fixture is only useful while the stored list is absent") + require.Nil(t, stored.DNSLabels) + + changed, err := stored.WouldChange(ConfigInput{NATExternalIPs: make([]string, 0)}) + require.NoError(t, err) + require.False(t, changed, "emptying a NAT list the profile never had is not a change") + + changed, err = stored.WouldChange(ConfigInput{DNSLabels: domain.List{}}) + require.NoError(t, err) + require.False(t, changed, "emptying a DNS label list the profile never had is not a change") + + // A list that does hold something still moves when the request empties it. + withEntries, err := UpdateConfig(ConfigInput{ConfigPath: path, NATExternalIPs: []string{"1.2.3.4"}}) + require.NoError(t, err) + require.Equal(t, []string{"1.2.3.4"}, withEntries.NATExternalIPs) + + changed, err = withEntries.WouldChange(ConfigInput{NATExternalIPs: make([]string, 0)}) + require.NoError(t, err) + require.True(t, changed, "clearing a NAT list that had an entry is a change") +} diff --git a/client/internal/profilemanager/invoking_user.go b/client/internal/profilemanager/invoking_user.go index c86a6ce43..7ba612ffb 100644 --- a/client/internal/profilemanager/invoking_user.go +++ b/client/internal/profilemanager/invoking_user.go @@ -6,6 +6,7 @@ import ( "os/user" "path/filepath" "runtime" + "strconv" log "github.com/sirupsen/logrus" ) @@ -13,17 +14,21 @@ import ( const envSudoUser = "SUDO_USER" var ( - geteuid = os.Geteuid - lookupUser = user.Lookup + currentUser = user.Current + getegid = os.Getegid + geteuid = os.Geteuid + lookupUser = user.Lookup ) // InvokingUser returns the user a CLI invocation acts for. Under sudo that is // the user who ran sudo, not root: privileged flags force commands through // sudo, and resolving profiles as root would silently switch the daemon to -// root's (default) profile instead of the invoking user's. Privilege decisions -// are not made here — those stay on the kernel credentials of the daemon -// connection, which SUDO_USER (a plain environment variable) can never -// influence; a forged value only selects a profile root could select anyway. +// root's (default) profile instead of the invoking user's. An unmapped positive +// process UID uses its numeric kernel identity; root, sudo lookup failures, and +// unavailable platform identities still fail closed. Privilege decisions stay +// on the kernel credentials of the daemon connection, which SUDO_USER (a plain +// environment variable) can never influence; a forged value only selects a +// profile root could select anyway. func InvokingUser() (*user.User, error) { if u, ok := sudoInvokingUser(); ok { return u, nil @@ -35,7 +40,23 @@ func InvokingUser() (*user.User, error) { if sudoActive() { return nil, fmt.Errorf("resolve sudo invoking user %q: refusing to fall back to root", os.Getenv(envSudoUser)) } - return user.Current() + u, err := currentUser() + if err == nil { + return u, nil + } + + uid := geteuid() + if uid <= 0 { + return nil, err + } + + log.Debugf("current user lookup for UID %d: %v; using numeric UID", uid, err) + uidString := strconv.Itoa(uid) + return &user.User{ + Username: uidString, + Uid: uidString, + Gid: strconv.Itoa(getegid()), + }, nil } // IsPlainRoot reports that the process runs as root with no usable sudo diff --git a/client/internal/profilemanager/invoking_user_test.go b/client/internal/profilemanager/invoking_user_test.go index 54c8ad8fd..159d2616b 100644 --- a/client/internal/profilemanager/invoking_user_test.go +++ b/client/internal/profilemanager/invoking_user_test.go @@ -2,6 +2,7 @@ package profilemanager import ( "errors" + "fmt" "io/fs" "os" "os/user" @@ -21,7 +22,51 @@ func TestInvokingUserFallsBackToProcessUser(t *testing.T) { current, err := user.Current() require.NoError(t, err) - assert.Equal(t, current.Username, got.Username) + assert.Equal(t, current.Username, got.Username, "invoking user should match the process user without sudo") +} + +func TestInvokingUserFailsClosedWithoutPositiveUID(t *testing.T) { + for _, uid := range []int{0, -1} { + t.Run(fmt.Sprintf("UID%d", uid), func(t *testing.T) { + t.Setenv(envSudoUser, "") + lookupErr := errors.New("current user unavailable") + fakeUnmappedUser(t, uid, 0, lookupErr) + + got, err := InvokingUser() + require.ErrorIs(t, err, lookupErr) + assert.Nil(t, got, "root or unavailable UID must not become a synthetic identity") + }) + } +} + +func TestProfileFilePathUsesNumericIdentityForUnmappedNonRoot(t *testing.T) { + t.Setenv(envSudoUser, "") + fakeUnmappedUser(t, 1001230000, 0, errors.New("user: unknown userid 1001230000")) + + profilesRoot := t.TempDir() + origDir := DefaultConfigPathDir + origOverride := ConfigDirOverride + DefaultConfigPathDir = profilesRoot + ConfigDirOverride = "" + t.Cleanup(func() { + DefaultConfigPathDir = origDir + ConfigDirOverride = origOverride + }) + + profileID := ID("0123456789abcdef0123456789abcdef") + got, err := (&Profile{ID: profileID}).FilePath() + require.NoError(t, err) + assert.Equal(t, + filepath.Join(profilesRoot, "1001230000", profileID.String()+".json"), + got, + "profile path should use the numeric UID namespace", + ) + + entries, err := os.ReadDir(profilesRoot) + require.NoError(t, err) + require.Len(t, entries, 1, "only the numeric UID directory should be created") + assert.Equal(t, "1001230000", entries[0].Name(), "profile namespace should be numeric") + assert.True(t, entries[0].IsDir(), "profile namespace should be a directory") } func TestSudoInvokingUserInactiveWithoutSudoContext(t *testing.T) { @@ -60,6 +105,13 @@ func TestInvokingUserFailsClosedWhenSudoLookupFails(t *testing.T) { fakeSudo(t, filepath.Join("/home", "misha")) lookupUser = func(string) (*user.User, error) { return nil, errors.New("nss unavailable") } + origCurrentUser := currentUser + currentUser = func() (*user.User, error) { + t.Fatal("currentUser must not be called after a sudo lookup failure") + return nil, errors.New("currentUser called unexpectedly") + } + t.Cleanup(func() { currentUser = origCurrentUser }) + got, err := InvokingUser() require.Error(t, err) assert.Nil(t, got, "must not resolve to the root process user") @@ -215,6 +267,22 @@ func fakeSudo(t *testing.T, home string) { }) } +func fakeUnmappedUser(t *testing.T, uid, gid int, lookupErr error) { + t.Helper() + + origCurrentUser := currentUser + origEuid := geteuid + origEgid := getegid + currentUser = func() (*user.User, error) { return nil, lookupErr } + geteuid = func() int { return uid } + getegid = func() int { return gid } + t.Cleanup(func() { + currentUser = origCurrentUser + geteuid = origEuid + getegid = origEgid + }) +} + func assertNoEntries(t *testing.T, root string) { t.Helper() err := filepath.WalkDir(root, func(path string, _ fs.DirEntry, err error) error { diff --git a/client/internal/profilemanager/service.go b/client/internal/profilemanager/service.go index ec287f01a..e58f421fd 100644 --- a/client/internal/profilemanager/service.go +++ b/client/internal/profilemanager/service.go @@ -313,7 +313,11 @@ func (s *ServiceManager) AddProfile(displayName, username string) (*Profile, err } profPath := filepath.Join(configDir, id.String()+".json") - cfg, err := createNewConfig(ConfigInput{ConfigPath: profPath}) + // Provisioned, not bare: this config goes straight to disk, and a profile + // file with no identity is one whose first reader has to mint the keys and + // remember to write them back. Before identity generation moved out of + // apply() into EnsureIdentity, createNewConfig produced them here too. + cfg, err := createProvisionedConfig(ConfigInput{ConfigPath: profPath}) if err != nil { return nil, fmt.Errorf("failed to create new config: %w", err) } @@ -330,6 +334,19 @@ func (s *ServiceManager) AddProfile(displayName, username string) (*Profile, err }, nil } +// RenameProfile changes a profile's display name. It rewrites the whole +// profile file, not just the name: the config is read through the normalizing +// reader, so apply()'s resolved values — the optional booleans, the interface +// blacklist, the DNS route interval — are persisted along with the new name. +// +// That is deliberate. A write that skipped apply() is what left profiles on +// disk carrying null where a value was meant, and made a diff of the config +// compare presence instead of value. Two consequences worth knowing: the +// platform-dependent defaults resolved here are the renaming host's +// (ServerSSHAllowed and the network monitor differ per OS), and a profile +// whose stored name does not survive sanitizeDisplayName now fails to rename +// rather than being rewritten — though apply() rejects such a profile on every +// other read too, so it was already unusable. func (s *ServiceManager) RenameProfile(id ID, username string, newName string) error { displayName, err := sanitizeDisplayName(newName) if err != nil { @@ -356,17 +373,17 @@ func (s *ServiceManager) RenameProfile(id ID, username string, newName string) e return ErrProfileNotFound } - data, err := os.ReadFile(target.Path) + // Through the reader, not a bare Unmarshal: this was the one write that + // skipped apply(), so it copied back whatever the file held — including an + // optional field left unset, which every other write resolves to its + // default. Renaming a profile is a poor place to leave that behind. + cfg, err := GetExistingConfig(target.Path) if err != nil { - return err - } - var cfg Config - if err := json.Unmarshal(data, &cfg); err != nil { - return err + return fmt.Errorf("read profile config: %w", err) } cfg.Name = displayName - if err := util.WriteJson(context.Background(), target.Path, cfg); err != nil { + if err := WriteOutConfig(target.Path, cfg); err != nil { return fmt.Errorf("failed to write profile name: %w", err) } return nil diff --git a/client/internal/profilemanager/service_test.go b/client/internal/profilemanager/service_test.go index 5e051b15d..d26ce746a 100644 --- a/client/internal/profilemanager/service_test.go +++ b/client/internal/profilemanager/service_test.go @@ -228,3 +228,27 @@ func TestRemoveProfile_DeletesStateFile(t *testing.T) { assert.True(t, errors.Is(err, os.ErrNotExist), "state file should be removed") }) } + +// A profile file is written here and read back by whoever connects with it, so +// it has to carry the peer's identity. While AddProfile used the bare +// constructor, it wrote a config with no keys: the first reader had to mint +// them, and the paths that read without writing — a gate deciding whether to +// refuse a request, the mobile SDKs loading a stored profile — got a config +// that cannot connect. +func TestAddProfileWritesAnIdentity(t *testing.T) { + withTestSM(t, func(sm *ServiceManager, username string) { + created, err := sm.AddProfile("work", username) + require.NoError(t, err) + + stored, err := GetExistingConfig(created.Path) + require.NoError(t, err) + + require.NotEmpty(t, stored.PrivateKey, "the profile was written without a WireGuard key") + require.NotEmpty(t, stored.SSHKey, "the profile was written without an SSH key") + + // And the identity is the one on disk, not one minted per read. + reread, err := GetExistingConfig(created.Path) + require.NoError(t, err) + require.Equal(t, stored.PrivateKey, reread.PrivateKey) + }) +} diff --git a/client/internal/relay/relay.go b/client/internal/relay/relay.go index 051717608..e39eed39a 100644 --- a/client/internal/relay/relay.go +++ b/client/internal/relay/relay.go @@ -201,11 +201,7 @@ func (p *StunTurnProbe) probeSTUN(ctx context.Context, uri *stun.URI) (addr stri } }() - net, err := stdnet.NewNet(ctx, nil) - if err != nil { - probeErr = fmt.Errorf("new net: %w", err) - return - } + net := stdnet.NewNet(ctx, nil, nil) client, err := stun.DialURI(uri, &stun.DialConfig{ Net: net, @@ -290,11 +286,7 @@ func (p *StunTurnProbe) probeTURN(ctx context.Context, uri *stun.URI) (addr stri } }() - net, err := stdnet.NewNet(ctx, nil) - if err != nil { - probeErr = fmt.Errorf("new net: %w", err) - return - } + net := stdnet.NewNet(ctx, nil, nil) cfg := &turn.ClientConfig{ STUNServerAddr: turnServerAddr, TURNServerAddr: turnServerAddr, diff --git a/client/internal/routemanager/client/client.go b/client/internal/routemanager/client/client.go index c691c54f8..973cf1ab8 100644 --- a/client/internal/routemanager/client/client.go +++ b/client/internal/routemanager/client/client.go @@ -294,6 +294,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error { return fmt.Errorf("add allowed IPs for peer %s: %w", route.Peer, err) } + w.statusRecorder.AddActiveRoutePeer(route.GetHAUniqueID(), route.Peer) if err := w.statusRecorder.AddPeerStateRoute(route.Peer, w.handler.String(), route.GetResourceID()); err != nil { log.Warnf("Failed to update peer state: %v", err) } @@ -303,6 +304,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error { } func (w *Watcher) removeAllowedIPs(route *route.Route, rsn reason) error { + w.statusRecorder.RemoveActiveRoutePeer(route.GetHAUniqueID()) if err := w.statusRecorder.RemovePeerStateRoute(route.Peer, w.handler.String()); err != nil { log.Warnf("Failed to update peer state: %v", err) } diff --git a/client/internal/routemanager/ipfwdstate/ipfwdstate.go b/client/internal/routemanager/ipfwdstate/ipfwdstate.go index 3d571e16b..22f7bd07a 100644 --- a/client/internal/routemanager/ipfwdstate/ipfwdstate.go +++ b/client/internal/routemanager/ipfwdstate/ipfwdstate.go @@ -19,8 +19,7 @@ type IPForwardingState struct { // routingV4/routingV6 track whether the routing path currently holds a // reference, so repeated EnableRouting calls (one per network-map update) - // hold at most one reference per family and an unpaired DisableRouting - // can't release references held by DNAT rules. + // hold at most one reference per family. routingV4 bool routingV6 bool @@ -95,31 +94,6 @@ func (f *IPForwardingState) ReleaseRouting() error { return nil } -// RequestForwarding enables the family's forwarding sysctl on first request. -func (f *IPForwardingState) RequestForwarding(v6 bool) error { - f.mu.Lock() - defer f.mu.Unlock() - - if v6 { - return f.requestV6() - } - return f.requestV4() -} - -// ReleaseForwarding decrements the family counter. The last v6 release restores -// what enable captured. v4 stays on: net.ipv4.ip_forward is co-owned by other -// tooling (docker, k8s, libvirt). -func (f *IPForwardingState) ReleaseForwarding(v6 bool) error { - f.mu.Lock() - defer f.mu.Unlock() - - if v6 { - return f.releaseV6() - } - f.releaseV4() - return nil -} - func (f *IPForwardingState) requestV4() error { if f.v4Count == 0 { if err := systemops.EnableV4IPForwarding(); err != nil { diff --git a/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go b/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go index b4615ff02..75209965c 100644 --- a/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go +++ b/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go @@ -10,8 +10,7 @@ import ( ) // TestRequestRoutingV6ToV4Transition verifies that a v4-only routing request -// releases a previously held routing-owned v6 reference without touching -// references held by DNAT rules. +// releases a previously held routing-owned v6 reference. func TestRequestRoutingV6ToV4Transition(t *testing.T) { f := NewIPForwardingState("wt-fwd-test") @@ -25,13 +24,6 @@ func TestRequestRoutingV6ToV4Transition(t *testing.T) { assert.Equal(t, 1, v4, "v4 reference kept") assert.Equal(t, 0, v6, "routing-owned v6 reference released") - // A DNAT-held reference survives a v4-only routing request. - require.NoError(t, f.RequestForwarding(true), "dnat v6 reference") - require.NoError(t, f.RequestRouting(false), "repeat v4-only request") - _, v6 = f.Counts() - assert.Equal(t, 1, v6, "dnat-held v6 reference survives") - require.NoError(t, f.ReleaseForwarding(true), "release dnat v6 reference") - require.NoError(t, f.ReleaseRouting(), "release routing") v4, v6 = f.Counts() assert.Equal(t, 0, v4, "all v4 references released") diff --git a/client/internal/routemanager/manager_test.go b/client/internal/routemanager/manager_test.go index 18b44820a..766ea1c61 100644 --- a/client/internal/routemanager/manager_test.go +++ b/client/internal/routemanager/manager_test.go @@ -8,6 +8,7 @@ import ( "net/netip" "testing" + "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/internal/stdnet" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" @@ -406,10 +407,7 @@ func TestManagerUpdateRoutes(t *testing.T) { for n, testCase := range testCases { t.Run(testCase.name, func(t *testing.T) { peerPrivateKey, _ := wgtypes.GeneratePrivateKey() - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil) opts := iface.WGIFaceOpts{ IFaceName: fmt.Sprintf("utun43%d", n), Address: wgaddr.MustParseWGAddress("100.65.65.2/24"), diff --git a/client/internal/routemanager/systemops/systemops_generic_test.go b/client/internal/routemanager/systemops/systemops_generic_test.go index c4f739c30..f117ff751 100644 --- a/client/internal/routemanager/systemops/systemops_generic_test.go +++ b/client/internal/routemanager/systemops/systemops_generic_test.go @@ -15,6 +15,7 @@ import ( "syscall" "testing" + "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/internal/stdnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -436,8 +437,7 @@ func createWGInterface(t *testing.T, interfaceName, ipAddressCIDR string, listen peerPrivateKey, err := wgtypes.GeneratePrivateKey() require.NoError(t, err) - newNet, err := stdnet.NewNet(context.Background(), nil) - require.NoError(t, err) + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist, nil) opts := iface.WGIFaceOpts{ IFaceName: interfaceName, diff --git a/client/internal/stdnet/filter.go b/client/internal/stdnet/filter.go index e45714001..07025ff1e 100644 --- a/client/internal/stdnet/filter.go +++ b/client/internal/stdnet/filter.go @@ -3,17 +3,13 @@ package stdnet import ( "runtime" "strings" - - log "github.com/sirupsen/logrus" - "golang.zx2c4.com/wireguard/wgctrl" ) // InterfaceFilter is a function passed to ICE Agent to filter out not allowed interfaces -// to avoid building tunnel over them. -func InterfaceFilter(disallowList []string) func(string) bool { - +// to avoid building tunnel over them. A nil detector probes the interface on every call, +// which is what the callers that build one filter for their whole lifetime want. +func InterfaceFilter(disallowList []string, detector *WGDetector) func(string) bool { return func(iFace string) bool { - if strings.HasPrefix(iFace, "lo") { // hardcoded loopback check to support already installed agents return false @@ -24,17 +20,8 @@ func InterfaceFilter(disallowList []string) func(string) bool { return false } } - // look for unlisted WireGuard interfaces - wg, err := wgctrl.New() - if err != nil { - log.Debugf("trying to create a wgctrl client failed with: %v", err) - return true - } - defer func() { - _ = wg.Close() - }() - _, err = wg.Device(iFace) - return err != nil + // look for unlisted WireGuard interfaces + return !detector.IsWireGuard(iFace) } } diff --git a/client/internal/stdnet/stdnet.go b/client/internal/stdnet/stdnet.go index 381886ac6..32f030612 100644 --- a/client/internal/stdnet/stdnet.go +++ b/client/internal/stdnet/stdnet.go @@ -45,12 +45,12 @@ type Net struct { } // NewNetWithDiscover creates a new StdNet instance. -func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) (*Net, error) { +func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string, detector *WGDetector) *Net { if ctx == nil { ctx = context.Background() } n := &Net{ - interfaceFilter: InterfaceFilter(disallowList), + interfaceFilter: InterfaceFilter(disallowList, detector), ctx: ctx, } // current ExternalIFaceDiscover implement in android-client https://github.dev/netbirdio/android-client @@ -60,20 +60,19 @@ func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover } else { n.iFaceDiscover = newMobileIFaceDiscover(iFaceDiscover) } - return n, n.UpdateInterfaces() + return n } // NewNet creates a new StdNet instance. -func NewNet(ctx context.Context, disallowList []string) (*Net, error) { +func NewNet(ctx context.Context, disallowList []string, detector *WGDetector) *Net { if ctx == nil { ctx = context.Background() } - n := &Net{ + return &Net{ iFaceDiscover: pionDiscover{}, - interfaceFilter: InterfaceFilter(disallowList), + interfaceFilter: InterfaceFilter(disallowList, detector), ctx: ctx, } - return n, n.UpdateInterfaces() } // resolveAddr performs DNS resolution with context support and timeout. @@ -122,45 +121,18 @@ func (n *Net) resolveAddr(network, address string) (netip.AddrPort, error) { return netip.AddrPortFrom(addrs[0], uint16(port)), nil } -// UpdateInterfaces updates the internal list of network interfaces -// and associated addresses filtering them by name. -// The interfaces are discovered by an external iFaceDiscover function or by a default discoverer if the external one -// wasn't specified. -func (n *Net) UpdateInterfaces() (err error) { - n.mu.Lock() - defer n.mu.Unlock() - - return n.updateInterfaces() -} - -func (n *Net) updateInterfaces() (err error) { - allIfaces, err := n.iFaceDiscover.iFaces() - if err != nil { - return err - } - - n.interfaces = n.filterInterfaces(allIfaces) - - n.lastUpdate = time.Now() - - return nil -} - // Interfaces returns a slice of interfaces which are available on the // system func (n *Net) Interfaces() ([]*transport.Interface, error) { n.mu.Lock() defer n.mu.Unlock() - if time.Since(n.lastUpdate) < updateInterval { - return slices.Clone(n.interfaces), nil + iFaces, err := n.freshInterfacesLocked() + if err != nil { + return nil, err } - if err := n.updateInterfaces(); err != nil { - return nil, fmt.Errorf("update interfaces: %w", err) - } - - return slices.Clone(n.interfaces), nil + return slices.Clone(iFaces), nil } // InterfaceByIndex returns the interface specified by index. @@ -171,7 +143,13 @@ func (n *Net) Interfaces() ([]*transport.Interface, error) { func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) { n.mu.Lock() defer n.mu.Unlock() - for _, ifc := range n.interfaces { + + iFaces, err := n.freshInterfacesLocked() + if err != nil { + return nil, err + } + + for _, ifc := range iFaces { if ifc.Index == index { return ifc, nil } @@ -184,7 +162,13 @@ func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) { func (n *Net) InterfaceByName(name string) (*transport.Interface, error) { n.mu.Lock() defer n.mu.Unlock() - for _, ifc := range n.interfaces { + + iFaces, err := n.freshInterfacesLocked() + if err != nil { + return nil, err + } + + for _, ifc := range iFaces { if ifc.Name == name { return ifc, nil } @@ -193,6 +177,31 @@ func (n *Net) InterfaceByName(name string) (*transport.Interface, error) { return nil, fmt.Errorf("%w: %s", transport.ErrInterfaceNotFound, name) } +func (n *Net) freshInterfacesLocked() ([]*transport.Interface, error) { + if time.Since(n.lastUpdate) < updateInterval { + return n.interfaces, nil + } + + if err := n.updateInterfacesLocked(); err != nil { + return nil, fmt.Errorf("update interfaces: %w", err) + } + + return n.interfaces, nil +} + +func (n *Net) updateInterfacesLocked() error { + allIFaces, err := n.iFaceDiscover.iFaces() + if err != nil { + return err + } + + n.interfaces = n.filterInterfaces(allIFaces) + + n.lastUpdate = time.Now() + + return nil +} + func (n *Net) filterInterfaces(interfaces []*transport.Interface) []*transport.Interface { if n.interfaceFilter == nil { return interfaces diff --git a/client/internal/stdnet/stdnet_test.go b/client/internal/stdnet/stdnet_test.go new file mode 100644 index 000000000..2b16a50c0 --- /dev/null +++ b/client/internal/stdnet/stdnet_test.go @@ -0,0 +1,136 @@ +package stdnet + +import ( + "context" + "errors" + "net" + "testing" + + "github.com/pion/transport/v3" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type countingDiscover struct { + calls int + list []*transport.Interface + err error +} + +func (d *countingDiscover) iFaces() ([]*transport.Interface, error) { + d.calls++ + if d.err != nil { + return nil, d.err + } + return d.list, nil +} + +func newTestNet(t *testing.T, d iFaceDiscover) *Net { + t.Helper() + return &Net{ + iFaceDiscover: d, + ctx: context.Background(), + } +} + +func testIFace(index int, name string) *transport.Interface { + return transport.NewInterface(net.Interface{Index: index, Name: name}) +} + +func TestNet_InterfacesDiscoversLazilyAndCaches(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}} + n := newTestNet(t, d) + + require.Zero(t, d.calls, "construction must not discover interfaces") + + iFaces, err := n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + assert.Equal(t, 1, d.calls) + + _, err = n.Interfaces() + require.NoError(t, err) + assert.Equal(t, 1, d.calls) +} + +func TestNewNet_DoesNotDiscoverAtConstruction(t *testing.T) { + n := NewNet(context.Background(), nil, nil) + require.NotNil(t, n) + assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold") +} + +func TestNewNetWithDiscover_DoesNotDiscoverAtConstruction(t *testing.T) { + n := NewNetWithDiscover(context.Background(), nil, nil, nil) + require.NotNil(t, n) + assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold") +} + +func TestNet_InterfacesRetryAfterDiscoveryFailure(t *testing.T) { + discoverErr := errors.New("discover failed") + d := &countingDiscover{err: discoverErr} + n := newTestNet(t, d) + + _, err := n.Interfaces() + require.ErrorIs(t, err, discoverErr) + + d.err = nil + d.list = []*transport.Interface{testIFace(1, "eth0")} + + iFaces, err := n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + assert.Equal(t, 2, d.calls) +} + +func TestNet_InterfaceByNameRefreshes(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}} + n := newTestNet(t, d) + + ifc, err := n.InterfaceByName("eth0") + require.NoError(t, err) + assert.Equal(t, "eth0", ifc.Name) + assert.Equal(t, 1, d.calls) + + _, err = n.InterfaceByName("nope") + require.ErrorIs(t, err, transport.ErrInterfaceNotFound) +} + +func TestNet_InterfaceByIndexRefreshes(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}} + n := newTestNet(t, d) + + ifc, err := n.InterfaceByIndex(3) + require.NoError(t, err) + assert.Equal(t, "eth0", ifc.Name) + assert.Equal(t, 1, d.calls) + + _, err = n.InterfaceByIndex(99) + require.ErrorIs(t, err, transport.ErrInterfaceNotFound) +} + +func TestNet_InterfaceLookupPropagatesDiscoveryError(t *testing.T) { + discoverErr := errors.New("discover failed") + n := newTestNet(t, &countingDiscover{err: discoverErr}) + + _, err := n.InterfaceByName("eth0") + require.ErrorIs(t, err, discoverErr) + + _, err = n.InterfaceByIndex(1) + require.ErrorIs(t, err, discoverErr) +} + +func TestNet_InterfacesReturnsCopy(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}} + n := newTestNet(t, d) + + iFaces, err := n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + + iFaces[0] = testIFace(2, "tampered") + + iFaces, err = n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + assert.Equal(t, "eth0", iFaces[0].Name) +} diff --git a/client/internal/stdnet/wgdetector.go b/client/internal/stdnet/wgdetector.go new file mode 100644 index 000000000..f34e469e1 --- /dev/null +++ b/client/internal/stdnet/wgdetector.go @@ -0,0 +1,119 @@ +package stdnet + +import ( + "sync" + "time" + + log "github.com/sirupsen/logrus" + "golang.org/x/sync/singleflight" + "golang.zx2c4.com/wireguard/wgctrl" +) + +// wgDetectorTTL bounds how long a cached answer is trusted. An interface rarely +// becomes, or stops being, a WireGuard device, and the window only has to be short +// enough that ICE does not keep gathering candidates on one that just appeared. +const wgDetectorTTL = 1 * time.Second + +type wgDetectorEntry struct { + isWireGuard bool + expireAt time.Time +} + +// WGDetector answers whether an interface is a WireGuard device, remembering the +// answer for a short while. +// +// The question is asked once per interface for every ICE agent, and an agent is +// created per peer connection attempt, so on a large network the uncached form runs +// constantly. Answering it means opening a wgctrl client, which builds both a kernel +// and a userspace client and resolves the netlink family, and then a round trip that +// usually just reports the device does not exist. +// +// A detector is safe for concurrent use and is meant to be shared by every agent. +type WGDetector struct { + ttl time.Duration + // probe is replaced in tests; it is the call this type exists to avoid repeating. + probe func(string) bool + + mu sync.RWMutex + cache map[string]wgDetectorEntry + + sf singleflight.Group +} + +// NewWGDetector returns a detector with the default time to live. +func NewWGDetector() *WGDetector { + return &WGDetector{ + ttl: wgDetectorTTL, + probe: probeWireGuard, + cache: make(map[string]wgDetectorEntry), + } +} + +// IsWireGuard reports whether the named interface is a WireGuard device. An interface +// it cannot ask about is reported as not being one, which leaves it available to ICE +// exactly as an uncached probe would. +func (d *WGDetector) IsWireGuard(iFace string) bool { + if d == nil { + return probeWireGuard(iFace) + } + + if isWireGuard, ok := d.cached(iFace); ok { + return isWireGuard + } + + result, _, _ := d.sf.Do(iFace, func() (interface{}, error) { + // A caller that saw the entry expire may get here after another caller already + // refreshed it and left the singleflight group. + if isWireGuard, ok := d.cached(iFace); ok { + return isWireGuard, nil + } + + isWireGuard := d.probe(iFace) + + d.store(iFace, isWireGuard) + + return isWireGuard, nil + }) + return result.(bool) +} + +func (d *WGDetector) cached(iFace string) (isWireGuard, ok bool) { + d.mu.RLock() + defer d.mu.RUnlock() + + entry, found := d.cache[iFace] + if !found || !time.Now().Before(entry.expireAt) { + return false, false + } + return entry.isWireGuard, true +} + +// store records an answer and drops the expired ones, so names of interfaces that came +// and went, such as container veths, do not pile up for the life of the engine. +func (d *WGDetector) store(iFace string, isWireGuard bool) { + now := time.Now() + + d.mu.Lock() + defer d.mu.Unlock() + + for name, entry := range d.cache { + if !now.Before(entry.expireAt) { + delete(d.cache, name) + } + } + d.cache[iFace] = wgDetectorEntry{isWireGuard: isWireGuard, expireAt: now.Add(d.ttl)} +} + +func probeWireGuard(iFace string) bool { + wg, err := wgctrl.New() + if err != nil { + log.Debugf("trying to create a wgctrl client failed with: %v", err) + return false + } + defer func() { + _ = wg.Close() + }() + + _, err = wg.Device(iFace) + return err == nil +} diff --git a/client/internal/stdnet/wgdetector_test.go b/client/internal/stdnet/wgdetector_test.go new file mode 100644 index 000000000..68260ae59 --- /dev/null +++ b/client/internal/stdnet/wgdetector_test.go @@ -0,0 +1,152 @@ +package stdnet + +import ( + "runtime" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// newCountingDetector returns a detector whose probe records how often it ran, so a test +// can assert on the thing this type exists for rather than on its return value alone. +func newCountingDetector(t *testing.T, ttl time.Duration, answer bool) (*WGDetector, *atomic.Int64) { + t.Helper() + + var calls atomic.Int64 + d := &WGDetector{ + ttl: ttl, + cache: make(map[string]wgDetectorEntry), + probe: func(string) bool { + calls.Add(1) + return answer + }, + } + return d, &calls +} + +func TestWGDetectorAsksOncePerInterfaceWithinTheTTL(t *testing.T) { + d, calls := newCountingDetector(t, time.Minute, true) + + for i := 0; i < 20; i++ { + assert.True(t, d.IsWireGuard("wt0"), "cached answer must not change") + } + assert.Equal(t, int64(1), calls.Load(), "the interface must be probed once within the TTL") + + d.IsWireGuard("eth0") + assert.Equal(t, int64(2), calls.Load(), "a different interface is a different question and is probed on its own") +} + +func TestWGDetectorReprobesAfterTheTTL(t *testing.T) { + d, calls := newCountingDetector(t, time.Millisecond, true) + + require.True(t, d.IsWireGuard("wt0"), "first answer") + require.Equal(t, int64(1), calls.Load(), "first call probes") + + time.Sleep(5 * time.Millisecond) + + require.True(t, d.IsWireGuard("wt0"), "answer after expiry") + assert.Equal(t, int64(2), calls.Load(), "an expired entry must be probed again") +} + +func TestWGDetectorCollapsesConcurrentProbes(t *testing.T) { + var calls atomic.Int64 + release := make(chan struct{}) + d := &WGDetector{ + ttl: time.Minute, + cache: make(map[string]wgDetectorEntry), + probe: func(string) bool { + calls.Add(1) + <-release + return true + }, + } + + var wg sync.WaitGroup + for i := 0; i < 50; i++ { + wg.Add(1) + go func() { + defer wg.Done() + d.IsWireGuard("wt0") + }() + } + + // The sleep only lets the callers pile up on the blocked probe so the collapse is + // exercised. The count does not depend on it: a caller that arrives after the probe + // finished finds the fresh entry, either before or inside the singleflight group. + time.Sleep(20 * time.Millisecond) + close(release) + wg.Wait() + + assert.Equal(t, int64(1), calls.Load(), "concurrent callers must share one probe") +} + +func TestWGDetectorNilProbesEveryTime(t *testing.T) { + var d *WGDetector + // A nil detector keeps the uncached behaviour, which is what the callers that build one + // filter for their whole lifetime rely on. It must not panic. + assert.NotPanics(t, func() { d.IsWireGuard("definitely-not-an-interface-0") }, + "a nil detector must fall back to probing") +} + +func TestInterfaceFilter(t *testing.T) { + wgDetector, calls := newCountingDetector(t, time.Minute, true) + plainDetector, _ := newCountingDetector(t, time.Minute, false) + + t.Run("loopback is rejected without probing", func(t *testing.T) { + filter := InterfaceFilter(nil, wgDetector) + assert.False(t, filter("lo"), "loopback must never be offered to ICE") + assert.Equal(t, int64(0), calls.Load(), "a name settled by prefix must not reach the probe") + }) + + t.Run("a disallowed interface is rejected without probing", func(t *testing.T) { + if runtime.GOOS == "ios" { + t.Skip("the disallow list is not applied on iOS") + } + filter := InterfaceFilter([]string{"wt"}, wgDetector) + assert.False(t, filter("wt0"), "a blacklisted interface must be rejected") + assert.Equal(t, int64(0), calls.Load(), "a name settled by the disallow list must not reach the probe") + }) + + t.Run("an unlisted WireGuard interface is rejected", func(t *testing.T) { + filter := InterfaceFilter(nil, wgDetector) + assert.False(t, filter("somewg0"), "a WireGuard interface must not be used to build a tunnel") + }) + + t.Run("an ordinary interface is allowed", func(t *testing.T) { + filter := InterfaceFilter(nil, plainDetector) + assert.True(t, filter("eth0"), "a plain interface must remain available to ICE") + }) +} + +func TestInterfaceFilterSharesOneProbeAcrossFilters(t *testing.T) { + d, calls := newCountingDetector(t, time.Minute, false) + + // Every ICE agent builds its own filter, twice, and each one is asked about every + // interface. Sharing the detector is what keeps that from repeating the probe. + for i := 0; i < 10; i++ { + filter := InterfaceFilter(nil, d) + require.True(t, filter("eth0"), "a plain interface stays allowed") + require.True(t, filter("eth1"), "a plain interface stays allowed") + } + + assert.Equal(t, int64(2), calls.Load(), "one probe per interface, not per filter") +} + +func TestWGDetectorDropsExpiredEntries(t *testing.T) { + d, _ := newCountingDetector(t, time.Millisecond, false) + + for _, name := range []string{"veth1", "veth2", "veth3"} { + d.IsWireGuard(name) + } + time.Sleep(5 * time.Millisecond) + d.IsWireGuard("eth0") + + d.mu.RLock() + defer d.mu.RUnlock() + assert.Len(t, d.cache, 1, "expired entries of vanished interfaces must be dropped") + assert.Contains(t, d.cache, "eth0", "the fresh answer must be kept") +} diff --git a/client/internal/updater/manager.go b/client/internal/updater/manager.go index 1b69368d0..c730713d3 100644 --- a/client/internal/updater/manager.go +++ b/client/internal/updater/manager.go @@ -21,8 +21,16 @@ const ( latestVersion = "latest" ) +const ( + modeUndecided updateMode = iota + modeDownloadOnly + modeManaged +) + var errNoUpdateState = errors.New("no update state found") +type updateMode int + type UpdateState struct { PreUpdateVersion string TargetVersion string @@ -36,8 +44,9 @@ type Manager struct { statusRecorder *peer.Status stateManager *statemanager.Manager - downloadOnly bool // true when no enforcement from management; notifies UI to download latest - forceUpdate bool // true when management sets AlwaysUpdate; skips UI interaction and installs directly + mode updateMode + modeGen uint64 + forceUpdate bool // true when management sets AlwaysUpdate; skips UI interaction and installs directly lastTrigger time.Time mgmUpdateChan chan struct{} @@ -54,7 +63,7 @@ type Manager struct { pendingVersion *v.Version // updateMutex protects update, expectedVersion, updateToLatestVersion, - // downloadOnly, forceUpdate, pendingVersion, and lastTrigger fields + // mode, modeGen, forceUpdate, pendingVersion, and lastTrigger fields updateMutex sync.Mutex // installMutex and installing guard against concurrent installation attempts @@ -76,7 +85,6 @@ func NewManager(statusRecorder *peer.Status, stateManager *statemanager.Manager) updateChannel: make(chan struct{}, 1), currentVersion: version.NetbirdVersion(), update: version.NewUpdate("nb/client"), - downloadOnly: true, autoUpdateSupported: isAutoUpdateSupported, } @@ -151,11 +159,7 @@ func (m *Manager) Start(ctx context.Context) { func (m *Manager) SetDownloadOnly() { m.updateMutex.Lock() - m.downloadOnly = true - m.forceUpdate = false - m.expectedVersion = nil - m.updateToLatestVersion = false - m.lastTrigger = time.Time{} + m.setModeLocked(modeDownloadOnly) m.updateMutex.Unlock() select { @@ -169,6 +173,7 @@ func (m *Manager) SetVersion(expectedVersion string, forceUpdate bool) { if !m.autoUpdateSupported() { log.Warnf("auto-update not supported on this platform") + m.SetDownloadOnly() return } @@ -177,30 +182,32 @@ func (m *Manager) SetVersion(expectedVersion string, forceUpdate bool) { if expectedVersion == "" { log.Errorf("empty expected version provided") - m.expectedVersion = nil - m.updateToLatestVersion = false - m.downloadOnly = true + m.setModeLocked(modeDownloadOnly) return } - if expectedVersion == latestVersion { - m.updateToLatestVersion = true - m.expectedVersion = nil - } else { - expectedSemVer, err := v.NewVersion(expectedVersion) + var expectedSemVer *v.Version + if expectedVersion != latestVersion { + parsed, err := v.NewVersion(expectedVersion) if err != nil { - log.Errorf("error parsing version: %v", err) + log.Errorf("error parsing version, switching to download-only: %v", err) + m.setModeLocked(modeDownloadOnly) + select { + case m.mgmUpdateChan <- struct{}{}: + default: + } return } - if m.expectedVersion != nil && m.expectedVersion.Equal(expectedSemVer) { - return - } - m.expectedVersion = expectedSemVer - m.updateToLatestVersion = false + expectedSemVer = parsed } - m.lastTrigger = time.Time{} - m.downloadOnly = false + if m.sameDirectiveLocked(expectedSemVer, forceUpdate) { + return + } + + m.setModeLocked(modeManaged) + m.expectedVersion = expectedSemVer + m.updateToLatestVersion = expectedSemVer == nil m.forceUpdate = forceUpdate select { @@ -209,6 +216,13 @@ func (m *Manager) SetVersion(expectedVersion string, forceUpdate bool) { } } +func (m *Manager) ResetMode() { + m.updateMutex.Lock() + defer m.updateMutex.Unlock() + + m.setModeLocked(modeUndecided) +} + // Install triggers the installation of the pending version. It is called when the user clicks the install button in the UI. func (m *Manager) Install(ctx context.Context) error { if !m.autoUpdateSupported() { @@ -255,12 +269,17 @@ func (m *Manager) NotifyUI() { m.updateMutex.Unlock() return } - downloadOnly := m.downloadOnly + mode := m.mode + gen := m.modeGen pendingVersion := m.pendingVersion latestVersion := m.update.LatestVersion() m.updateMutex.Unlock() - if downloadOnly { + if mode == modeUndecided { + return + } + + if mode == modeDownloadOnly { if latestVersion == nil { return } @@ -268,6 +287,9 @@ func (m *Manager) NotifyUI() { if err != nil || currentVersion.GreaterThanOrEqual(latestVersion) { return } + if m.modeChanged(gen) { + return + } m.statusRecorder.PublishEvent( cProto.SystemEvent_INFO, cProto.SystemEvent_SYSTEM, @@ -278,7 +300,7 @@ func (m *Manager) NotifyUI() { return } - if pendingVersion != nil { + if pendingVersion != nil && !m.modeChanged(gen) { m.statusRecorder.PublishEvent( cProto.SystemEvent_INFO, cProto.SystemEvent_SYSTEM, @@ -343,13 +365,18 @@ func (m *Manager) handleUpdate(ctx context.Context) { return } - downloadOnly := m.downloadOnly + mode := m.mode + gen := m.modeGen forceUpdate := m.forceUpdate curLatestVersion := m.update.LatestVersion() switch { + case mode == modeUndecided: + log.Tracef("auto-update mode not decided yet") + m.updateMutex.Unlock() + return // Download-only mode or resolve "latest" to actual version - case downloadOnly, m.updateToLatestVersion: + case mode == modeDownloadOnly, m.updateToLatestVersion: if curLatestVersion == nil { log.Tracef("latest version not fetched yet") m.updateMutex.Unlock() @@ -374,12 +401,17 @@ func (m *Manager) handleUpdate(ctx context.Context) { m.lastTrigger = time.Now() log.Infof("new version available: %s", updateVersion) - if !downloadOnly && !forceUpdate { + if mode == modeManaged && !forceUpdate { m.pendingVersion = updateVersion } m.updateMutex.Unlock() - if downloadOnly { + if m.modeChanged(gen) { + log.Debugf("auto-update mode changed while checking %s, discarding", updateVersion) + return + } + + if mode == modeDownloadOnly { m.statusRecorder.PublishEvent( cProto.SystemEvent_INFO, cProto.SystemEvent_SYSTEM, @@ -406,6 +438,33 @@ func (m *Manager) handleUpdate(ctx context.Context) { ) } +func (m *Manager) modeChanged(gen uint64) bool { + m.updateMutex.Lock() + defer m.updateMutex.Unlock() + + return m.modeGen != gen +} + +func (m *Manager) sameDirectiveLocked(expectedVersion *v.Version, forceUpdate bool) bool { + if m.mode != modeManaged || m.forceUpdate != forceUpdate { + return false + } + if expectedVersion == nil { + return m.updateToLatestVersion + } + return m.expectedVersion != nil && m.expectedVersion.Equal(expectedVersion) +} + +func (m *Manager) setModeLocked(mode updateMode) { + m.mode = mode + m.modeGen++ + m.forceUpdate = false + m.expectedVersion = nil + m.updateToLatestVersion = false + m.pendingVersion = nil + m.lastTrigger = time.Time{} +} + func (m *Manager) install(ctx context.Context, pendingVersion *v.Version) error { m.statusRecorder.PublishEvent( cProto.SystemEvent_CRITICAL, diff --git a/client/internal/updater/manager_linux_test.go b/client/internal/updater/manager_linux_test.go index b05dd7e7d..4501ddde7 100644 --- a/client/internal/updater/manager_linux_test.go +++ b/client/internal/updater/manager_linux_test.go @@ -16,7 +16,7 @@ import ( ) // On Linux, only Mode 1 (downloadOnly) is supported. -// SetVersion is a no-op because auto-update installation is not supported. +// SetVersion falls back to download-only because auto-update installation is not supported. func Test_LatestVersion_Linux(t *testing.T) { testMatrix := []struct { @@ -70,7 +70,7 @@ func Test_LatestVersion_Linux(t *testing.T) { t.Errorf("%s: Initial version mismatch, expected %v, got %v", c.name, c.initialLatestVersion.String(), ver) } - mockUpdate.latestVersion = c.latestVersion + mockUpdate.setLatestVersion(c.latestVersion) mockUpdate.onUpdate() ver, enforced = waitForUpdateEvent(sub, 500*time.Millisecond) @@ -89,22 +89,24 @@ func Test_LatestVersion_Linux(t *testing.T) { } } -func Test_SetVersion_NoOp_Linux(t *testing.T) { - // On Linux, SetVersion should be a no-op — no events fired - tmpFile := path.Join(t.TempDir(), "update-test-noop.json") +func Test_SetVersion_FallsBackToDownloadOnly_Linux(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-fallback.json") recorder := peer.NewRecorder("") sub := recorder.SubscribeToEvents() defer recorder.UnsubscribeFromEvents(sub) m := NewManager(recorder, statemanager.New(tmpFile)) - m.update = &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.1"))} + m.update = &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.5"))} m.currentVersion = "1.0.0" m.Start(context.Background()) m.SetVersion("1.0.1", false) - ver, _ := waitForUpdateEvent(sub, 500*time.Millisecond) - if ver != "" { - t.Errorf("SetVersion should be a no-op on Linux, but got event with version %s", ver) + ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.5" { + t.Fatalf("expected download-only event for fetched 1.0.5, got %q", ver) + } + if enforced { + t.Error("Linux fallback must never have enforced metadata") } m.Stop() diff --git a/client/internal/updater/manager_mode_test.go b/client/internal/updater/manager_mode_test.go new file mode 100644 index 000000000..aa238f94c --- /dev/null +++ b/client/internal/updater/manager_mode_test.go @@ -0,0 +1,217 @@ +package updater + +import ( + "context" + "path" + "testing" + "time" + + v "github.com/hashicorp/go-version" + + "github.com/netbirdio/netbird/client/internal/peer" + "github.com/netbirdio/netbird/client/internal/statemanager" +) + +func Test_UndecidedMode_SuppressesNotification(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-undecided.json") + recorder := peer.NewRecorder("") + sub := recorder.SubscribeToEvents() + defer recorder.UnsubscribeFromEvents(sub) + + mockUpdate := &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.1"))} + m := NewManager(recorder, statemanager.New(tmpFile)) + m.update = mockUpdate + m.currentVersion = "1.0.0" + m.Start(context.Background()) + defer m.Stop() + + mockUpdate.onUpdate() + if ver, _ := waitForUpdateEvent(sub, 300*time.Millisecond); ver != "" { + t.Fatalf("undecided mode must not publish, got %q", ver) + } + + m.NotifyUI() + if ver, _ := waitForUpdateEvent(sub, 300*time.Millisecond); ver != "" { + t.Fatalf("NotifyUI in undecided mode must not publish, got %q", ver) + } + + m.SetDownloadOnly() + ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.1" { + t.Fatalf("expected download-only event for 1.0.1, got %q", ver) + } + if enforced { + t.Error("download-only event must not carry enforced metadata") + } +} + +func Test_ResetMode_ReturnsToUndecided(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-reset.json") + recorder := peer.NewRecorder("") + sub := recorder.SubscribeToEvents() + defer recorder.UnsubscribeFromEvents(sub) + + mockUpdate := &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.1"))} + m := NewManager(recorder, statemanager.New(tmpFile)) + m.update = mockUpdate + m.currentVersion = "1.0.0" + m.autoUpdateSupported = func() bool { return true } + m.Start(context.Background()) + defer m.Stop() + + m.SetVersion("1.0.1", false) + ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.1" || !enforced { + t.Fatalf("expected enforced event for 1.0.1, got %q enforced=%v", ver, enforced) + } + + m.ResetMode() + + mockUpdate.onUpdate() + if ver, _ := waitForUpdateEvent(sub, 300*time.Millisecond); ver != "" { + t.Fatalf("reset mode must not publish on fetch, got %q", ver) + } + + m.NotifyUI() + if ver, _ := waitForUpdateEvent(sub, 300*time.Millisecond); ver != "" { + t.Fatalf("NotifyUI after reset must not publish, got %q", ver) + } + + if err := m.Install(context.Background()); err == nil { + t.Fatal("Install after reset must fail without a pending version") + } + + m.SetVersion("1.0.1", false) + ver, enforced = waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.1" || !enforced { + t.Fatalf("expected enforced event again after reset, got %q enforced=%v", ver, enforced) + } +} + +func Test_SetDownloadOnly_ClearsPendingVersion(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-pending.json") + recorder := peer.NewRecorder("") + sub := recorder.SubscribeToEvents() + defer recorder.UnsubscribeFromEvents(sub) + + m := NewManager(recorder, statemanager.New(tmpFile)) + m.update = &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.1"))} + m.currentVersion = "1.0.0" + m.autoUpdateSupported = func() bool { return true } + m.Start(context.Background()) + defer m.Stop() + + m.SetVersion("1.0.1", false) + if ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond); ver != "1.0.1" || !enforced { + t.Fatalf("expected enforced event for 1.0.1, got %q enforced=%v", ver, enforced) + } + + m.SetDownloadOnly() + if ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond); ver != "1.0.1" || enforced { + t.Fatalf("expected download-only event for 1.0.1, got %q enforced=%v", ver, enforced) + } + + if err := m.Install(context.Background()); err == nil { + t.Fatal("Install in download-only mode must not install the staged managed version") + } +} + +func Test_ResetMode_SilencesStaleForceDirective(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-force-reset.json") + recorder := peer.NewRecorder("") + sub := recorder.SubscribeToEvents() + defer recorder.UnsubscribeFromEvents(sub) + + mockUpdate := &versionUpdateMock{} + m := NewManager(recorder, statemanager.New(tmpFile)) + m.update = mockUpdate + m.currentVersion = "1.0.0" + m.autoUpdateSupported = func() bool { return true } + m.Start(context.Background()) + defer m.Stop() + + // Management enforces "latest" before the fetcher has reported any version, + // so nothing can be installed while the engine is still up. + m.SetVersion(latestVersion, true) + if event := waitForAnyEvent(sub, 300*time.Millisecond); event != nil { + t.Fatalf("no event expected before the latest version is known, got %v", event) + } + + // The engine stop resets the mode. A release published afterwards must not + // trigger the stale forced install or any notification. + m.ResetMode() + mockUpdate.setLatestVersion(v.Must(v.NewSemver("1.0.1"))) + mockUpdate.onUpdate() + if event := waitForAnyEvent(sub, 300*time.Millisecond); event != nil { + t.Fatalf("stale force directive must stay silent after reset, got %v", event) + } + + m.SetVersion("1.0.1", false) + ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.1" || !enforced { + t.Fatalf("expected enforced event after a fresh directive, got %q enforced=%v", ver, enforced) + } +} + +func Test_SetVersion_MalformedFallsBackToDownloadOnly(t *testing.T) { + tmpFile := path.Join(t.TempDir(), "update-test-malformed.json") + recorder := peer.NewRecorder("") + sub := recorder.SubscribeToEvents() + defer recorder.UnsubscribeFromEvents(sub) + + m := NewManager(recorder, statemanager.New(tmpFile)) + m.update = &versionUpdateMock{latestVersion: v.Must(v.NewSemver("1.0.1"))} + m.currentVersion = "1.0.0" + m.autoUpdateSupported = func() bool { return true } + m.Start(context.Background()) + defer m.Stop() + + m.SetVersion("not-a-version", false) + ver, enforced := waitForUpdateEvent(sub, 500*time.Millisecond) + if ver != "1.0.1" { + t.Fatalf("expected download-only event for 1.0.1 after malformed version, got %q", ver) + } + if enforced { + t.Error("malformed version fallback must not carry enforced metadata") + } +} + +func Test_SetVersion_ForceChangeAppliesWithSameVersion(t *testing.T) { + m := NewManager(peer.NewRecorder(""), statemanager.New(path.Join(t.TempDir(), "update-test-force-change.json"))) + m.update = &versionUpdateMock{} + m.autoUpdateSupported = func() bool { return true } + + m.SetVersion("1.0.1", false) + m.SetVersion("1.0.1", true) + + m.updateMutex.Lock() + defer m.updateMutex.Unlock() + if !m.forceUpdate { + t.Fatal("enabling force update without a version change must take effect") + } + if m.expectedVersion == nil || m.expectedVersion.String() != "1.0.1" { + t.Fatalf("expected version 1.0.1 to stay set, got %v", m.expectedVersion) + } +} + +func Test_SetVersion_RepeatedDirectiveKeepsMode(t *testing.T) { + m := NewManager(peer.NewRecorder(""), statemanager.New(path.Join(t.TempDir(), "update-test-repeat.json"))) + m.update = &versionUpdateMock{} + m.autoUpdateSupported = func() bool { return true } + + for _, expected := range []string{"1.0.1", latestVersion} { + m.SetVersion(expected, false) + m.updateMutex.Lock() + gen := m.modeGen + m.updateMutex.Unlock() + + m.SetVersion(expected, false) + m.updateMutex.Lock() + repeatedGen := m.modeGen + m.updateMutex.Unlock() + + if repeatedGen != gen { + t.Errorf("repeating the %q directive must not reset the mode", expected) + } + } +} diff --git a/client/internal/updater/manager_test.go b/client/internal/updater/manager_test.go index 107dca2b3..939c09814 100644 --- a/client/internal/updater/manager_test.go +++ b/client/internal/updater/manager_test.go @@ -66,7 +66,7 @@ func Test_LatestVersion(t *testing.T) { t.Errorf("%s: Initial update version mismatch, expected %v, got %v", c.name, c.initialLatestVersion.String(), ver) } - mockUpdate.latestVersion = c.latestVersion + mockUpdate.setLatestVersion(c.latestVersion) mockUpdate.onUpdate() ver, _ = waitForUpdateEvent(sub, 500*time.Millisecond) diff --git a/client/internal/updater/manager_test_helpers_test.go b/client/internal/updater/manager_test_helpers_test.go index c7faee1f4..430b31ec1 100644 --- a/client/internal/updater/manager_test_helpers_test.go +++ b/client/internal/updater/manager_test_helpers_test.go @@ -2,21 +2,24 @@ package updater import ( "strconv" + "sync" "time" v "github.com/hashicorp/go-version" "github.com/netbirdio/netbird/client/internal/peer" + cProto "github.com/netbirdio/netbird/client/proto" ) type versionUpdateMock struct { latestVersion *v.Version onUpdate func() + mu sync.Mutex } -func (m versionUpdateMock) StopWatch() {} +func (m *versionUpdateMock) StopWatch() {} -func (m versionUpdateMock) SetDaemonVersion(newVersion string) bool { +func (m *versionUpdateMock) SetDaemonVersion(newVersion string) bool { return false } @@ -24,11 +27,19 @@ func (m *versionUpdateMock) SetOnUpdateListener(updateFn func()) { m.onUpdate = updateFn } -func (m versionUpdateMock) LatestVersion() *v.Version { +func (m *versionUpdateMock) LatestVersion() *v.Version { + m.mu.Lock() + defer m.mu.Unlock() return m.latestVersion } -func (m versionUpdateMock) StartFetcher() {} +func (m *versionUpdateMock) StartFetcher() {} + +func (m *versionUpdateMock) setLatestVersion(version *v.Version) { + m.mu.Lock() + defer m.mu.Unlock() + m.latestVersion = version +} // waitForUpdateEvent waits for a new_version_available event, returns the version string or "" on timeout. func waitForUpdateEvent(sub *peer.EventSubscription, timeout time.Duration) (version string, enforced bool) { @@ -54,3 +65,20 @@ func waitForUpdateEvent(sub *peer.EventSubscription, timeout time.Duration) (ver } } } + +// waitForAnyEvent returns the first published event of any kind, or nil on timeout. +// Unlike waitForUpdateEvent it also catches the install-progress events, so a test +// can assert that a forced install never started. +func waitForAnyEvent(sub *peer.EventSubscription, timeout time.Duration) *cProto.SystemEvent { + timer := time.NewTimer(timeout) + defer timer.Stop() + select { + case event, ok := <-sub.Events(): + if !ok { + return nil + } + return event + case <-timer.C: + return nil + } +} diff --git a/client/internal/wincmd/system32_windows.go b/client/internal/wincmd/system32_windows.go new file mode 100644 index 000000000..36aa258b5 --- /dev/null +++ b/client/internal/wincmd/system32_windows.go @@ -0,0 +1,30 @@ +// Package wincmd locates the Windows utilities the client shells out to. +package wincmd + +import ( + "path/filepath" + + log "github.com/sirupsen/logrus" + "golang.org/x/sys/windows" +) + +// defaultSystem32Dir is where the system directory is on every supported +// install, used only when the API that reports it fails. +const defaultSystem32Dir = `C:\Windows\System32` + +// System32 returns the full path of a Windows utility under the system +// directory. +// +// PATH is deliberately not consulted. The daemon runs as LocalSystem with an +// environment of its own, so whoever can place an entry in that PATH chooses +// which binary runs with those privileges. The system directory is read from +// the API rather than from %SystemRoot% for the same reason. +func System32(command string) string { + sysDir, err := windows.GetSystemDirectory() + if err != nil { + log.Warnf("Failed to locate the Windows system directory, falling back to %s: %v", defaultSystem32Dir, err) + sysDir = defaultSystem32Dir + } + + return filepath.Join(sysDir, command+".exe") +} diff --git a/client/internal/wincmd/system32_windows_test.go b/client/internal/wincmd/system32_windows_test.go new file mode 100644 index 000000000..0d31d7ee7 --- /dev/null +++ b/client/internal/wincmd/system32_windows_test.go @@ -0,0 +1,31 @@ +package wincmd + +import ( + "os" + "path/filepath" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestSystem32IgnoresPATH(t *testing.T) { + // A directory holding something that would win a PATH lookup, in front of + // everything else: the daemon runs as LocalSystem, so a PATH entry must not + // be able to decide what it executes. + planted := t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(planted, "netsh.exe"), []byte("not really netsh"), 0o600)) + t.Setenv("PATH", planted+string(os.PathListSeparator)+os.Getenv("PATH")) + + got := System32("netsh") + + assert.True(t, filepath.IsAbs(got), "the path must be absolute, got %q", got) + assert.NotContains(t, got, planted, "a PATH entry must not be consulted") + assert.True(t, strings.EqualFold(filepath.Base(got), "netsh.exe"), "unexpected file name in %q", got) + + // The system directory is what Windows reports it to be, not %SystemRoot%, + // which the same caller could have set alongside PATH. + t.Setenv("SystemRoot", planted) + assert.Equal(t, got, System32("netsh"), "%SystemRoot% must not move the lookup") +} diff --git a/client/ios/NetBirdSDK/client.go b/client/ios/NetBirdSDK/client.go index bbbb969c9..457ae3a7d 100644 --- a/client/ios/NetBirdSDK/client.go +++ b/client/ios/NetBirdSDK/client.go @@ -88,9 +88,15 @@ type Client struct { // netMgr outlives engine restarts: it mirrors the OS connectivity, not // the engine lifecycle. Run injects its state and sweeper into each new // ConnectClient. - netMgr *netevents.Manager - // preloadedConfig holds config loaded from JSON (used on tvOS where file writes are blocked) - preloadedConfig *profilemanager.Config + netMgr *netevents.Manager + preloadedConfigJSON atomic.Pointer[string] + + // mdmSource holds the per-Client MDM policy source and its change + // detector as one unit. Set by SetMDMPolicyFetcher (called from the + // Swift side at extension init). 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] // stateMu guards the run lifecycle as one unit: the cancel installed by // the current run, the channel it closes on exit, and the state it @@ -122,44 +128,48 @@ func NewClient(cfgFile, stateFile, cacheDir, logFilePath, deviceName string, osV } } -// SetConfigFromJSON loads config from a JSON string into memory. -// This is used on tvOS where file writes to App Group containers are blocked. -// When set, IsLoginRequired() and Run() will use this preloaded config instead of reading from file. +// SetConfigFromJSON stores the JSON config that later loads resolve instead of the config file (tvOS). func (c *Client) SetConfigFromJSON(jsonStr string) error { - cfg, err := profilemanager.ConfigFromJSON(jsonStr) - if err != nil { + // Parsed only to reject an unreadable document early; the JSON itself is + // what is stored, and every load re-parses it. A document carrying no peer + // identity is readable and accepted: that is a logged-out profile, and the + // login that follows provisions the keys. + if _, err := profilemanager.ConfigFromJSON(jsonStr); err != nil { log.Errorf("SetConfigFromJSON: failed to parse config JSON: %v", err) return err } - c.preloadedConfig = cfg + c.preloadedConfigJSON.Store(&jsonStr) log.Infof("SetConfigFromJSON: config loaded successfully from JSON") return nil } +func (c *Client) loadConfig(input profilemanager.ConfigInput) (*profilemanager.Config, error) { + var cfg *profilemanager.Config + var err error + if preloaded := c.preloadedConfigJSON.Load(); preloaded != nil { + cfg, err = profilemanager.ConfigFromJSON(*preloaded) + } else { + cfg, err = profilemanager.DirectUpdateOrCreateConfig(input) + } + if err != nil { + return nil, err + } + c.applyMDMOverlay(cfg) + return cfg, nil +} + // Run start the internal client. It is a blocker function func (c *Client) Run(fd int32, interfaceName string, envList *EnvList) error { exportEnvList(envList) log.Infof("Starting NetBird client") log.Debugf("Tunnel uses interface: %s", interfaceName) - var cfg *profilemanager.Config - var err error - - // Use preloaded config if available (tvOS where file writes are blocked) - if c.preloadedConfig != nil { - log.Infof("Run: using preloaded config from memory") - cfg = c.preloadedConfig - } else { - log.Infof("Run: loading config from file") - // Use DirectUpdateOrCreateConfig to avoid atomic file operations (temp file + rename) - // which are blocked by the tvOS sandbox in App Group containers - cfg, err = profilemanager.DirectUpdateOrCreateConfig(profilemanager.ConfigInput{ - ConfigPath: c.cfgFile, - StateFilePath: c.stateFile, - }) - if err != nil { - return err - } + cfg, err := c.loadConfig(profilemanager.ConfigInput{ + ConfigPath: c.cfgFile, + StateFilePath: c.stateFile, + }) + if err != nil { + return err } c.recorder.UpdateManagementAddress(cfg.ManagementURL.String()) c.recorder.UpdateRosenpass(cfg.RosenpassEnabled, cfg.RosenpassPermissive) @@ -274,19 +284,13 @@ func (c *Client) DebugBundle(anonymize bool, anonymizeLevel string) (string, err // If the engine hasn't been started, load config so we can reach management. if cfg == nil { - if c.preloadedConfig != nil { - cfg = c.preloadedConfig - } else { - var err error - // Use DirectUpdateOrCreateConfig to avoid atomic file operations - // (temp file + rename) blocked by the tvOS sandbox. - cfg, err = profilemanager.DirectUpdateOrCreateConfig(profilemanager.ConfigInput{ - ConfigPath: c.cfgFile, - StateFilePath: c.stateFile, - }) - if err != nil { - return "", fmt.Errorf("load config: %w", err) - } + var err error + cfg, err = c.loadConfig(profilemanager.ConfigInput{ + ConfigPath: c.cfgFile, + StateFilePath: c.stateFile, + }) + if err != nil { + return "", fmt.Errorf("load config: %w", err) } } @@ -421,29 +425,9 @@ func (c *Client) IsLoginRequired() bool { ctx, cancel := context.WithCancel(ctxWithValues) defer cancel() - var cfg *profilemanager.Config - var err error - - // Use preloaded config if available (tvOS where file writes are blocked) - if c.preloadedConfig != nil { - log.Infof("IsLoginRequired: using preloaded config from memory") - cfg = c.preloadedConfig - } else { - log.Infof("IsLoginRequired: loading config from file") - // Use DirectUpdateOrCreateConfig to avoid atomic file operations (temp file + rename) - // which are blocked by the tvOS sandbox in App Group containers - cfg, err = profilemanager.DirectUpdateOrCreateConfig(profilemanager.ConfigInput{ - ConfigPath: c.cfgFile, - }) - if err != nil { - log.Errorf("IsLoginRequired: failed to load config: %v", err) - // If we can't load config, assume login is required - return true - } - } - - if cfg == nil { - log.Errorf("IsLoginRequired: config is nil") + cfg, err := c.loadConfig(profilemanager.ConfigInput{ConfigPath: c.cfgFile}) + if err != nil { + log.Errorf("IsLoginRequired: failed to load config: %v", err) return true } @@ -493,8 +477,9 @@ func (c *Client) LoginForMobile() string { log.Errorf("LoginForMobile: failed to load config: %v", err) return fmt.Sprintf("failed to load config: %v", err) } + c.applyMDMOverlay(cfg) - oAuthFlow, err := auth.NewOAuthFlow(ctx, cfg, false, false, "") + oAuthFlow, err := auth.NewOAuthFlow(ctx, cfg, false, false, "", false) if err != nil { return err.Error() } diff --git a/client/ios/NetBirdSDK/login.go b/client/ios/NetBirdSDK/login.go index cf7aa6730..b182e82b6 100644 --- a/client/ios/NetBirdSDK/login.go +++ b/client/ios/NetBirdSDK/login.go @@ -11,6 +11,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" ) @@ -39,14 +40,22 @@ type Auth struct { ctx context.Context cancel context.CancelFunc config *profilemanager.Config + base *profilemanager.Config + policy *mdm.Policy cfgPath string } -// NewAuth instantiate Auth struct and validate the management URL -func NewAuth(cfgPath string, mgmURL string) (*Auth, error) { - inputCfg := profilemanager.ConfigInput{ - ConfigPath: cfgPath, - ManagementURL: mgmURL, +// NewAuth instantiate Auth struct and validate the management URL. +// 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 } // Load the existing config when a config file is already present so an @@ -67,6 +76,10 @@ func NewAuth(cfgPath string, mgmURL string) (*Auth, error) { if err != nil { return nil, err } + a := &Auth{policy: policy, cfgPath: cfgPath} + if err := a.setBaseConfig(cfg); err != nil { + return nil, err + } // Use a cancellable context so Stop() can abort an in-progress interactive // login. The PKCE flow's WaitToken blocks (and keeps its loopback HTTP server @@ -76,14 +89,8 @@ func NewAuth(cfgPath string, mgmURL string) (*Auth, error) { // process (decoupled from the network extension), so without this the server // lingers after the user dismisses the browser and the next connect stalls // trying to bind the same port. - ctx, cancel := context.WithCancel(context.Background()) - - return &Auth{ - ctx: ctx, - cancel: cancel, - config: cfg, - cfgPath: cfgPath, - }, nil + a.ctx, a.cancel = context.WithCancel(context.Background()) + return a, nil } // NewAuthWithConfig instantiate Auth based on existing config @@ -106,9 +113,7 @@ func (a *Auth) Stop() { } } -// 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) { if listener == nil { log.Errorf("SaveConfigIfSSOSupported: listener is nil") @@ -136,17 +141,10 @@ func (a *Auth) saveConfigIfSSOSupported() (bool, error) { return false, fmt.Errorf("failed to check SSO support: %v", err) } - if !supportsSSO { - return false, nil - } - - // Use DirectWriteOutConfig to avoid atomic file operations (temp file + rename) - // which are blocked by the tvOS sandbox in App Group containers - err = profilemanager.DirectWriteOutConfig(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) { if resultListener == nil { log.Errorf("LoginWithSetupKeyAndSaveConfig: resultListener is nil") @@ -175,10 +173,7 @@ func (a *Auth) loginWithSetupKeyAndSaveConfig(setupKey string, deviceName string if err != nil { return fmt.Errorf("login failed: %v", err) } - - // Use DirectWriteOutConfig to avoid atomic file operations (temp file + rename) - // which are blocked by the tvOS sandbox in App Group containers - return profilemanager.DirectWriteOutConfig(a.cfgPath, a.config) + return nil } // LoginSync performs a synchronous login check without UI interaction @@ -312,19 +307,6 @@ func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName strin } } - // Save the config before notifying success to ensure persistence completes - // before the callback potentially triggers teardown on the Swift side. - // Note: This differs from Android which doesn't save config after login. - // On iOS/tvOS, we save here because: - // 1. The config may have been modified during login (e.g., new tokens) - // 2. On tvOS, the Network Extension context may be the only place with - // write permissions to the App Group container - if a.cfgPath != "" { - if err := profilemanager.DirectWriteOutConfig(a.cfgPath, a.config); err != nil { - log.Warnf("failed to save config after login: %v", err) - } - } - // Notify caller of successful login synchronously before returning urlOpener.OnLoginSuccess() @@ -348,7 +330,7 @@ func profileLoginHint(cfgPath string) string { const authInfoRequestTimeout = 30 * time.Second func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, forceDeviceAuth bool) (*auth.TokenInfo, error) { - oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, forceDeviceAuth, profileLoginHint(a.cfgPath)) + oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, forceDeviceAuth, false, profileLoginHint(a.cfgPath)) if err != nil { return nil, fmt.Errorf("failed to get OAuth flow: %v", err) } @@ -375,23 +357,73 @@ func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener return &tokenInfo, nil } -// GetConfigJSON returns the current config as a JSON string. -// This can be used by the caller to persist the config via alternative storage -// mechanisms (e.g., UserDefaults on tvOS where file writes are blocked). +// GetConfigJSON returns the config without the MDM overlay as JSON, for persisting it outside the config file (tvOS). func (a *Auth) GetConfigJSON() (string, error) { - if a.config == nil { + cfg := a.base + if cfg == nil { + cfg = a.config + } + if cfg == nil { return "", fmt.Errorf("no config available") } - return profilemanager.ConfigToJSON(a.config) + return profilemanager.ConfigToJSON(cfg) } -// SetConfigFromJSON loads config from a JSON string. -// This can be used to restore config from alternative storage mechanisms. +// SetConfigFromJSON replaces the config from JSON; the MDM overlay is applied on top for the login. func (a *Auth) SetConfigFromJSON(jsonStr string) error { cfg, err := profilemanager.ConfigFromJSON(jsonStr) if err != nil { return err } - a.config = cfg + return a.setBaseConfig(cfg) +} + +func (a *Auth) setBaseConfig(base *profilemanager.Config) error { + // A logged-out profile carries no keys: the mobile logout clears them in + // place so the next login registers a new peer instead of resurrecting the + // old one. This is that login, and auth.NewAuth parses the WireGuard key + // before the SSO flow even starts, so an absent identity fails the login on + // key size rather than asking the user to sign in. + // + // Minted on the base config, which is the one GetConfigJSON hands back for + // the caller to store — the overlaid copy below is runtime-only. + generated, err := base.EnsureIdentity() + if err != nil { + return fmt.Errorf("ensure profile identity: %w", err) + } + if generated { + if a.cfgPath != "" { + // Non-atomic, like NewAuth's own write: the tvOS App Group sandbox + // blocks the temp-file-and-rename an atomic write needs. + if err := profilemanager.DirectWriteOutConfig(a.cfgPath, base); err != nil { + return fmt.Errorf("write out profile config: %w", err) + } + } else { + // No file to write to — this is the tvOS path, where the profile + // lives in the caller's own store. It persists the new identity by + // calling GetConfigJSON once the login completes; until then the + // keys exist only here, and a login that never completes leaves + // nothing behind. + log.Infof("provisioned a peer identity for a config with no file on disk") + } + } + + overlaid, err := copyConfig(base) + if err != nil { + return err + } + if a.policy != nil { + overlaid.ApplyMDMPolicy(a.policy) + } + a.base = base + a.config = overlaid return nil } + +func copyConfig(cfg *profilemanager.Config) (*profilemanager.Config, error) { + raw, err := profilemanager.ConfigToJSON(cfg) + if err != nil { + return nil, err + } + return profilemanager.ConfigFromJSON(raw) +} diff --git a/client/ios/NetBirdSDK/mdm.go b/client/ios/NetBirdSDK/mdm.go new file mode 100644 index 000000000..93a31916c --- /dev/null +++ b/client/ios/NetBirdSDK/mdm.go @@ -0,0 +1,66 @@ +//go:build ios + +package NetBirdSDK + +import ( + "github.com/netbirdio/netbird/client/internal/profilemanager" + "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 +} + +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 +} + +func loaderFor(p PolicyFetcher) *mdm.Loader { + if p == nil { + return mdm.NewJSONLoader(nil) + } + return mdm.NewJSONLoader(p.FetchJSON) +} diff --git a/client/ios/NetBirdSDK/preferences.go b/client/ios/NetBirdSDK/preferences.go index 39aa7ed83..642f9e160 100644 --- a/client/ios/NetBirdSDK/preferences.go +++ b/client/ios/NetBirdSDK/preferences.go @@ -3,12 +3,16 @@ package NetBirdSDK import ( + "sync/atomic" + "github.com/netbirdio/netbird/client/internal/profilemanager" + "github.com/netbirdio/netbird/client/mdm" ) // Preferences export a subset of the internal config for gomobile type Preferences struct { configInput profilemanager.ConfigInput + mdmLoader atomic.Pointer[mdm.Loader] } // NewPreferences create new Preferences instance @@ -17,20 +21,39 @@ func NewPreferences(configPath string, stateFilePath string) *Preferences { ConfigPath: configPath, StateFilePath: stateFilePath, } - 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 read 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 store the given url and wait for commit @@ -44,7 +67,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 } @@ -56,17 +79,21 @@ func (p *Preferences) SetAdminURL(url string) { p.configInput.AdminURL = url } -// GetPreSharedKey read preshared 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 store the given key and wait for commit @@ -81,11 +108,14 @@ func (p *Preferences) SetRosenpassEnabled(enabled bool) { // GetRosenpassEnabled read rosenpass enabled 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 } @@ -99,11 +129,14 @@ func (p *Preferences) SetRosenpassPermissive(permissive bool) { // GetRosenpassPermissive read rosenpass permissive 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 } @@ -116,7 +149,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 } @@ -130,18 +163,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 @@ -151,6 +186,9 @@ func (p *Preferences) SetRemoteJobsAllowed(allowed bool) { // Commit write out the changes into config file func (p *Preferences) Commit() error { + if err := profilemanager.CheckMDMConflicts(p.configInput, p.policy()); err != nil { + return err + } // Use DirectUpdateOrCreateConfig to avoid atomic file operations (temp file + rename) // which are blocked by the tvOS sandbox in App Group containers _, err := profilemanager.DirectUpdateOrCreateConfig(p.configInput) diff --git a/client/ios/NetBirdSDK/preferences_test.go b/client/ios/NetBirdSDK/preferences_test.go index 5f75e7c9a..2382e123c 100644 --- a/client/ios/NetBirdSDK/preferences_test.go +++ b/client/ios/NetBirdSDK/preferences_test.go @@ -31,14 +31,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") } } @@ -69,13 +68,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") } } @@ -114,12 +113,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") } } diff --git a/client/ios/NetBirdSDK/profile_manager.go b/client/ios/NetBirdSDK/profile_manager.go index 139521c7f..df962e227 100644 --- a/client/ios/NetBirdSDK/profile_manager.go +++ b/client/ios/NetBirdSDK/profile_manager.go @@ -52,6 +52,12 @@ func NewProfileManager(configDir string) *ProfileManager { return &ProfileManager{impl: mobile.NewProfileManager(configDir, iosUsername)} } +// 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) { diff --git a/client/mdm/changedetector.go b/client/mdm/changedetector.go new file mode 100644 index 000000000..5c21ae355 --- /dev/null +++ b/client/mdm/changedetector.go @@ -0,0 +1,34 @@ +package mdm + +import "sync" + +// ChangeDetector tracks the last observed policy of a Loader so an +// OS-notification-driven caller can ask whether the managed configuration +// actually changed before restarting anything. +type ChangeDetector struct { + mu sync.Mutex + loader *Loader + prev *Policy +} + +// NewChangeDetector constructs a ChangeDetector seeded with the loader's +// current policy, so only a later change reports as changed. +func NewChangeDetector(loader *Loader) *ChangeDetector { + return &ChangeDetector{ + loader: loader, + prev: loader.Load(), + } +} + +// Changed re-reads the policy, logs the per-key diff, and reports whether it +// diverged from the last observation; the new snapshot becomes the baseline. +func (d *ChangeDetector) Changed() bool { + d.mu.Lock() + defer d.mu.Unlock() + curr := d.loader.Load() + if !policyChanged(d.prev, curr) { + return false + } + d.prev = curr + return true +} diff --git a/client/mdm/conflicts.go b/client/mdm/conflicts.go new file mode 100644 index 000000000..160212afb --- /dev/null +++ b/client/mdm/conflicts.go @@ -0,0 +1,116 @@ +package mdm + +import ( + "net/url" + + "github.com/netbirdio/netbird/util" +) + +// PreSharedKeyRedactedSentinel is the redaction mask returned in place of a +// real pre-shared key; an incoming value equal to it is a round-trip echo, +// never an override. +const PreSharedKeyRedactedSentinel = "**********" + +// ConflictCheck is a value-aware comparison between a single requested field +// and the corresponding MDM-enforced value. +type ConflictCheck struct { + Key string + Check func(*Policy) bool +} + +// ConflictBool builds a ConflictCheck for a boolean MDM key. +func ConflictBool(key string, p *bool) ConflictCheck { + return ConflictCheck{ + Key: key, + Check: func(pol *Policy) bool { + if p == nil { + return true + } + want, ok := pol.GetBool(key) + return ok && want == *p + }, + } +} + +// ConflictStringPtr builds a ConflictCheck for an optional string MDM key, +// where an explicit empty value is still a request to change the setting. A +// nil p means "field not set" (no override requested). +func ConflictStringPtr(key string, p *string) ConflictCheck { + return ConflictCheck{ + Key: key, + Check: func(pol *Policy) bool { + if p == nil { + return true + } + want, ok := pol.GetString(key) + return ok && want == *p + }, + } +} + +// ConflictURL builds a ConflictCheck for a URL-typed MDM key. The two sides are +// compared as the endpoints they address, not as strings: see +// util.SameServiceURL. +func ConflictURL(key, got string) ConflictCheck { + return ConflictCheck{ + Key: key, + Check: func(pol *Policy) bool { + if got == "" { + return true + } + want, ok := pol.GetString(key) + return ok && util.SameServiceURLStrings(want, got) + }, + } +} + +// ConflictInt64 builds a ConflictCheck for an integer MDM key. +func ConflictInt64(key string, p *int64) ConflictCheck { + return ConflictCheck{ + Key: key, + Check: func(pol *Policy) bool { + if p == nil { + return true + } + want, ok := pol.GetInt(key) + return ok && want == *p + }, + } +} + +// ResolveConflicts returns the names of keys whose requested value diverges +// from the policy-enforced value; keys the policy does not manage are skipped, +// a managed key without a Check counts as a conflict. +func ResolveConflicts(policy *Policy, checks []ConflictCheck) []string { + if policy.IsEmpty() { + return nil + } + var conflicts []string + for _, c := range checks { + if !policy.HasKey(c.Key) { + continue + } + if c.Check == nil || !c.Check(policy) { + conflicts = append(conflicts, c.Key) + } + } + return conflicts +} + +// CanonicalURL normalizes a service URL by appending the scheme default port +// when none is present; unparseable input is returned unchanged. +func CanonicalURL(s string) string { + u, err := url.ParseRequestURI(s) + if err != nil { + return s + } + if u.Port() == "" { + switch u.Scheme { + case "https": + u.Host += ":443" + case "http": + u.Host += ":80" + } + } + return u.String() +} diff --git a/client/mdm/conflicts_test.go b/client/mdm/conflicts_test.go new file mode 100644 index 000000000..d145ec103 --- /dev/null +++ b/client/mdm/conflicts_test.go @@ -0,0 +1,40 @@ +package mdm + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// The same spellings, through the conflict check that decides whether a request +// is refused. An enforced URL restated in another spelling addresses the very +// server the policy names, so it must not be reported as a conflict. +func TestConflictURLComparesEndpoints(t *testing.T) { + policy := NewPolicy(map[string]any{KeyManagementURL: "https://mgmt.example.com"}) + require.True(t, policy.HasKey(KeyManagementURL)) + + for _, restated := range []string{ + "https://mgmt.example.com", + "https://mgmt.example.com:443", + "https://mgmt.example.com/", + "https://MGMT.example.com", + "https://mgmt.example.com:0443", + } { + conflicts := ResolveConflicts(policy, []ConflictCheck{ConflictURL(KeyManagementURL, restated)}) + assert.Empty(t, conflicts, "%q is the enforced endpoint written differently", restated) + } + + for _, diverging := range []string{ + "https://other.example.com", + "http://mgmt.example.com", + "https://mgmt.example.com:8443", + "https://mgmt.example.com/other", + } { + conflicts := ResolveConflicts(policy, []ConflictCheck{ConflictURL(KeyManagementURL, diverging)}) + assert.Equal(t, []string{KeyManagementURL}, conflicts, "%q addresses another endpoint", diverging) + } + + // An unset field is not a request to change anything. + assert.Empty(t, ResolveConflicts(policy, []ConflictCheck{ConflictURL(KeyManagementURL, "")})) +} diff --git a/client/mdm/jsonloader.go b/client/mdm/jsonloader.go new file mode 100644 index 000000000..7139b0e4f --- /dev/null +++ b/client/mdm/jsonloader.go @@ -0,0 +1,34 @@ +package mdm + +import ( + "encoding/json" + + log "github.com/sirupsen/logrus" +) + +type jsonPolicyFetcher struct { + fetch func() string +} + +// NewJSONLoader constructs a Loader whose policy source is a JSON-encoded +// object string, as produced by the mobile native layers; a nil fetch +// disables MDM enforcement. +func NewJSONLoader(fetch func() string) *Loader { + if fetch == nil { + return NewLoader(nil) + } + return NewLoader(&jsonPolicyFetcher{fetch: fetch}) +} + +func (f *jsonPolicyFetcher) Fetch() map[string]any { + raw := f.fetch() + if raw == "" { + return nil + } + var out map[string]any + if err := json.Unmarshal([]byte(raw), &out); err != nil { + log.Warnf("MDM mobile fetcher: invalid JSON payload from native: %v", err) + return nil + } + return out +} diff --git a/client/mdm/policy.go b/client/mdm/policy.go index dac135ea6..638fa0d80 100644 --- a/client/mdm/policy.go +++ b/client/mdm/policy.go @@ -119,16 +119,46 @@ func NewPolicy(values map[string]any) *Policy { return &Policy{values: values} } -// LoadPolicy reads the platform-native MDM configuration. Returns an -// empty (but non-nil) Policy when no source is present, the source is -// empty, or the platform is unsupported. +// PolicyFetcher supplies the managed configuration to a Loader. Mobile +// platforms (Android / iOS) implement it to push the OS-managed values +// into the Go runtime. On every platform a non-nil fetcher takes +// precedence over the native source, which is the test seam for the +// registry / plist loaders; a nil fetcher leaves the native source in +// charge, or disables MDM enforcement where there is none. +type PolicyFetcher interface { + Fetch() map[string]any +} + +// Loader is the DI-friendly entry point for reading the active MDM +// policy. Construct one at the daemon's lifecycle owner (Server on +// desktop, gomobile-exposed bridge on mobile) and pass it to anything +// that needs to read MDM state (the reload ticker, profilemanager's +// Config). Each callsite has the Loader handed in instead of looking +// up package-level state. +type Loader struct { + fetcher PolicyFetcher +} + +// NewLoader constructs a Loader. A non-nil fetcher takes precedence over +// the platform-native source; production desktop callers pass nil so the +// registry / plist stays authoritative. +func NewLoader(f PolicyFetcher) *Loader { + return &Loader{fetcher: f} +} + +// Load reads the platform-native MDM configuration and returns a +// Policy. Returns an empty (but non-nil) Policy when no source is +// present, the source is empty, or the platform is unsupported. // // Diagnostic logging differentiates the three states: // - source absent / unsupported platform: trace log only // - source present, zero keys: info "MDM enrolled (no managed keys)" // - source present, N keys: info "MDM enrolled with N managed keys: [...]" -func LoadPolicy() *Policy { - values, err := loadPlatformPolicy() +func (l *Loader) Load() *Policy { + if l == nil { + return &Policy{values: map[string]any{}} + } + values, err := l.loadPlatform() if err != nil { log.Tracef("MDM policy load: %v", err) return &Policy{values: map[string]any{}} @@ -205,6 +235,8 @@ func (p *Policy) GetBool(key string) (bool, bool) { return t != 0, true case int64: return t != 0, true + case float64: + return t != 0, true } return false, false } @@ -270,7 +302,7 @@ func (p *Policy) GetStringSlice(key string) ([]string, bool) { } // sortedKeys returns the keys of m as a deterministic, lexicographically -// sorted slice. Used internally by Policy.ManagedKeys and LoadPolicy's +// sorted slice. Used internally by Policy.ManagedKeys and Loader.Load's // diagnostic log line so callers see a stable key order across runs // regardless of Go's randomised map iteration. func sortedKeys(m map[string]any) []string { diff --git a/client/mdm/policy_darwin.go b/client/mdm/policy_darwin.go index 57aa1168c..4159f5b7e 100644 --- a/client/mdm/policy_darwin.go +++ b/client/mdm/policy_darwin.go @@ -25,8 +25,9 @@ import ( // writable plist, as a defense against tampered installs. const policyPlistPath = "/Library/Managed Preferences/io.netbird.client.plist" -// loadPlatformPolicy reads the MDM-managed configuration from the macOS -// managed-preferences plist at policyPlistPath. Returns: +// loadPlatform reads the MDM-managed configuration from the macOS +// managed-preferences plist at policyPlistPath, unless a fetcher was +// injected, in which case its values are returned instead. Returns: // - (nil, nil) when the plist is absent (device not MDM-enrolled for // NetBird, or admin has not yet pushed a payload) // - (map, nil) with N entries when N managed values are present @@ -39,13 +40,19 @@ const policyPlistPath = "/Library/Managed Preferences/io.netbird.client.plist" // skipped so a stray entry in the payload does not block startup. // Native plist value types map naturally onto the Policy accessor // expectations (GetString / GetBool / GetInt / GetStringSlice). -func loadPlatformPolicy() (map[string]any, error) { +func (l *Loader) loadPlatform() (map[string]any, error) { + // Honour the injected fetcher when present so tests (and any + // future non-macOS MDM channel) can short-circuit the plist read + // with a scripted policy. + if l != nil && l.fetcher != nil { + return l.fetcher.Fetch(), nil + } f, err := os.Open(policyPlistPath) if err != nil { if errors.Is(err, fs.ErrNotExist) { // Not enrolled for NetBird. Caller treats nil as // "no MDM source present". - //nolint:nilnil // (nil, nil) is the documented platform-absent sentinel; see LoadPolicy. + //nolint:nilnil // (nil, nil) is the documented platform-absent sentinel; see Loader.Load. return nil, nil } return nil, fmt.Errorf("open %s: %w", policyPlistPath, err) diff --git a/client/mdm/policy_mobile.go b/client/mdm/policy_mobile.go index ec25d4bb1..2e25a2bb5 100644 --- a/client/mdm/policy_mobile.go +++ b/client/mdm/policy_mobile.go @@ -2,13 +2,14 @@ package mdm -// loadPlatformPolicy is unused on mobile: the native layer (Swift on iOS, -// Kotlin/Java on Android) reads the OS managed-config store and pushes the -// resulting dictionary in-process via a gomobile entry point that lands in -// Phase 5 / Phase 6. The stub keeps the package compilable for mobile -// builds and returns (nil, nil) — the platform-absent sentinel that -// LoadPolicy in policy.go treats as "no MDM source present". -func loadPlatformPolicy() (map[string]any, error) { - //nolint:nilnil // (nil, nil) is the documented platform-absent sentinel; see LoadPolicy. - return nil, nil +// loadPlatform reads the OS-managed configuration via the native +// PolicyFetcher injected at Loader construction. Returns +// (nil, nil) — the platform-absent sentinel that Loader.Load treats as +// "no MDM source present" — when no fetcher was provided. +func (l *Loader) loadPlatform() (map[string]any, error) { + if l == nil || l.fetcher == nil { + //nolint:nilnil // (nil, nil) is the documented platform-absent sentinel; see Loader.Load. + return nil, nil + } + return l.fetcher.Fetch(), nil } diff --git a/client/mdm/policy_other.go b/client/mdm/policy_other.go index f4263afa2..5d0b17cfd 100644 --- a/client/mdm/policy_other.go +++ b/client/mdm/policy_other.go @@ -2,13 +2,17 @@ package mdm -// loadPlatformPolicy returns no policy on platforms without an MDM channel -// (Linux, FreeBSD). MDM enforcement is off and the client behaves as if -// the feature did not exist. Returns (nil, nil) — the platform-absent -// sentinel the caller (LoadPolicy in policy.go) treats as "no MDM -// source present"; an error here would just translate to the same -// outcome with an extra log line. -func loadPlatformPolicy() (map[string]any, error) { - //nolint:nilnil // (nil, nil) is the documented platform-absent sentinel; see LoadPolicy. +// loadPlatform reads the MDM policy on platforms without a native MDM +// channel (Linux, FreeBSD). When no fetcher was injected the policy is +// (nil, nil) — the platform-absent sentinel that Loader.Load treats as +// "MDM enforcement disabled". A non-nil fetcher takes precedence: it +// is the test-seam used by unit tests to inject a scripted policy +// without touching the OS, and the same hook supports any future +// non-mobile OS that grows an out-of-band MDM channel. +func (l *Loader) loadPlatform() (map[string]any, error) { + if l != nil && l.fetcher != nil { + return l.fetcher.Fetch(), nil + } + //nolint:nilnil // (nil, nil) is the documented platform-absent sentinel; see Loader.Load. return nil, nil } diff --git a/client/mdm/policy_test.go b/client/mdm/policy_test.go index 6cbe69776..ea467f861 100644 --- a/client/mdm/policy_test.go +++ b/client/mdm/policy_test.go @@ -1,6 +1,7 @@ package mdm import ( + "runtime" "testing" "github.com/stretchr/testify/assert" @@ -95,7 +96,8 @@ func TestPolicy_GetBool(t *testing.T) { {"int64 nonzero", int64(2), true, true}, {"int64 zero", int64(0), false, true}, {"string garbage", "maybe", false, false}, - {"float unsupported", 1.0, false, false}, + {"float nonzero", 1.0, true, true}, + {"float zero", 0.0, false, true}, } for _, c := range cases { t.Run(c.name, func(t *testing.T) { @@ -155,10 +157,29 @@ func TestPolicy_GetStringSlice(t *testing.T) { }) } -func TestLoadPolicy_PlatformStubReturnsEmpty(t *testing.T) { - // loadPlatformPolicy is a stub on every OS for Phase 1. LoadPolicy must - // degrade gracefully and never return nil. - p := LoadPolicy() +// encoding/json decodes every JSON number into float64, so the mobile +// loaders never see int. +func TestJSONLoader_BoolFromNumber(t *testing.T) { + p := NewJSONLoader(func() string { return `{"blockInbound":1,"disableProfiles":0}` }).Load() + + got, ok := p.GetBool(KeyBlockInbound) + assert.True(t, ok) + assert.True(t, got) + + got, ok = p.GetBool(KeyDisableProfiles) + assert.True(t, ok) + assert.False(t, got) +} + +func TestLoader_NilFetcherReturnsEmpty(t *testing.T) { + // Loader.Load with no fetcher (desktop construction) must degrade + // gracefully and never return nil; on linux loadPlatform is a stub + // returning (nil, nil), and Load is expected to translate that + // into a non-nil empty Policy. + if runtime.GOOS == "windows" || runtime.GOOS == "darwin" { + t.Skip("a nil fetcher reads the OS-managed policy on this platform") + } + p := NewLoader(nil).Load() require.NotNil(t, p) assert.True(t, p.IsEmpty()) assert.Empty(t, p.ManagedKeys()) diff --git a/client/mdm/policy_windows.go b/client/mdm/policy_windows.go index 0c2629f98..9363db436 100644 --- a/client/mdm/policy_windows.go +++ b/client/mdm/policy_windows.go @@ -61,8 +61,9 @@ func readRegistryValue(k registry.Key, name, canonical string, out map[string]an } } -// loadPlatformPolicy reads the MDM-managed configuration from the -// Windows registry under HKLM\Software\Policies\NetBird. Returns: +// loadPlatform reads the MDM-managed configuration from the Windows +// registry under HKLM\Software\Policies\NetBird, unless a fetcher was +// injected, in which case its values are returned instead. Returns: // - (nil, nil) when the key is absent (device not MDM-enrolled for NetBird) // - (map, nil) with N entries when N managed values are set (N may be 0) // - (nil, err) on open / enumerate registry errors @@ -70,12 +71,18 @@ func readRegistryValue(k registry.Key, name, canonical string, out map[string]an // Per-value type coercion + skip-on-error is delegated to // readRegistryValue. Unknown value names are logged and skipped so a // malformed deployment does not block startup. -func loadPlatformPolicy() (map[string]any, error) { +func (l *Loader) loadPlatform() (map[string]any, error) { + // Honour the injected fetcher when present so tests (and any + // future non-Windows MDM channel) can short-circuit the registry + // read with a scripted policy. + if l != nil && l.fetcher != nil { + return l.fetcher.Fetch(), nil + } k, err := registry.OpenKey(registry.LOCAL_MACHINE, policyRegistryPath, registry.QUERY_VALUE) if err != nil { if errors.Is(err, registry.ErrNotExist) { // Not enrolled. Caller treats nil as "no MDM source present". - //nolint:nilnil // (nil, nil) is the documented platform-absent sentinel; see LoadPolicy. + //nolint:nilnil // (nil, nil) is the documented platform-absent sentinel; see Loader.Load. return nil, nil } return nil, fmt.Errorf("open %s: %w", policyRegistryPath, err) diff --git a/client/mdm/restrictions.go b/client/mdm/restrictions.go new file mode 100644 index 000000000..200756b78 --- /dev/null +++ b/client/mdm/restrictions.go @@ -0,0 +1,91 @@ +package mdm + +import "encoding/json" + +// Fields carries the per-key MDM enforcement state for a UI: value-typed +// fields hold the enforced value (nil pointer = not managed), boolean +// fields report that the key is managed. +type Fields struct { + ManagementURL string `json:"managementURL"` + PreSharedKey bool `json:"preSharedKey"` + WireguardPort bool `json:"wireguardPort"` + RosenpassEnabled bool `json:"rosenpassEnabled"` + RosenpassPermissive bool `json:"rosenpassPermissive"` + DisableClientRoutes bool `json:"disableClientRoutes"` + DisableServerRoutes bool `json:"disableServerRoutes"` + AllowServerSSH *bool `json:"allowServerSSH"` + DisableAutoConnect bool `json:"disableAutoConnect"` + DisableAutostart bool `json:"disableAutostart"` + BlockInbound bool `json:"blockInbound"` + DisableMetricsCollection bool `json:"disableMetricsCollection"` + SplitTunnelMode bool `json:"splitTunnelMode"` + SplitTunnelApps bool `json:"splitTunnelApps"` + RemoteJobsAllowed bool `json:"allowRemoteJobs"` + DisableAdvancedView *bool `json:"disableAdvancedView"` +} + +// Features carries the feature gates a UI must honor. +type Features struct { + DisableProfiles bool `json:"disableProfiles"` + DisableNetworks bool `json:"disableNetworks"` + DisableUpdateSettings bool `json:"disableUpdateSettings"` +} + +// Restrictions is the UI-facing enforcement snapshot; the JSON shape is +// shared by the desktop frontend and the mobile bridges. +type Restrictions struct { + MDM Fields `json:"mdm"` + Features Features `json:"features"` +} + +// BuildRestrictions derives the UI enforcement snapshot from the active +// policy. +func BuildRestrictions(policy *Policy) Restrictions { + var r Restrictions + if policy.IsEmpty() { + return r + } + + if v, ok := policy.GetString(KeyManagementURL); ok { + r.MDM.ManagementURL = CanonicalURL(v) + } + r.MDM.PreSharedKey = policy.HasKey(KeyPreSharedKey) + r.MDM.WireguardPort = policy.HasKey(KeyWireguardPort) + r.MDM.RosenpassEnabled = policy.HasKey(KeyRosenpassEnabled) + r.MDM.RosenpassPermissive = policy.HasKey(KeyRosenpassPermissive) + r.MDM.DisableClientRoutes = policy.HasKey(KeyDisableClientRoutes) + r.MDM.DisableServerRoutes = policy.HasKey(KeyDisableServerRoutes) + r.MDM.DisableAutoConnect = policy.HasKey(KeyDisableAutoConnect) + r.MDM.DisableAutostart = policy.HasKey(KeyDisableAutostart) + r.MDM.BlockInbound = policy.HasKey(KeyBlockInbound) + r.MDM.DisableMetricsCollection = policy.HasKey(KeyDisableMetricsCollection) + r.MDM.SplitTunnelMode = policy.HasKey(KeySplitTunnelMode) + r.MDM.SplitTunnelApps = policy.HasKey(KeySplitTunnelApps) + r.MDM.RemoteJobsAllowed = policy.HasKey(KeyRemoteJobsAllowed) + if v, ok := policy.GetBool(KeyAllowServerSSH); ok { + r.MDM.AllowServerSSH = &v + } + if v, ok := policy.GetBool(KeyDisableAdvancedView); ok { + r.MDM.DisableAdvancedView = &v + } + + if v, ok := policy.GetBool(KeyDisableProfiles); ok { + r.Features.DisableProfiles = v + } + if v, ok := policy.GetBool(KeyDisableNetworks); ok { + r.Features.DisableNetworks = v + } + if v, ok := policy.GetBool(KeyDisableUpdateSettings); ok { + r.Features.DisableUpdateSettings = v + } + return r +} + +// JSON renders the snapshot in the shared UI JSON shape. +func (r Restrictions) JSON() (string, error) { + b, err := json.Marshal(r) + if err != nil { + return "", err + } + return string(b), nil +} diff --git a/client/mdm/ticker.go b/client/mdm/ticker.go index abd6ae233..be8fdcce7 100644 --- a/client/mdm/ticker.go +++ b/client/mdm/ticker.go @@ -15,33 +15,33 @@ import ( // instead, hence anticipating the ticker mechanism entirely. const DefaultReloadInterval = 1 * time.Minute -// policyLoader is the indirection through which the ticker reads the -// OS-native policy, both for the initial observation and on every tick. -// Production points it at LoadPolicy; tests in this package override it to -// feed a scripted sequence of policies without touching the real OS store. -var policyLoader = LoadPolicy - -// Ticker periodically re-reads the OS-native MDM policy via LoadPolicy and -// invokes the onChange callback (supplied to Run) whenever the observed -// Policy diverges from the last observation (added / removed / changed -// keys). Launch with Run from a goroutine; cancel the supplied context -// to stop. +// Ticker periodically re-reads the OS-native MDM policy via the +// injected Loader and invokes the onChange callback (supplied to Run) +// whenever the observed Policy diverges from the last observation +// (added / removed / changed keys). Launch with Run from a goroutine; +// cancel the supplied context to stop. type Ticker struct { interval time.Duration + loader *Loader prev *Policy } // NewTicker constructs a Ticker that will re-read the OS-native policy -// every reloadInterval once Run is called. -// The initial snapshot is populated by calling policyLoader at +// every reloadInterval once Run is called. The Loader is injected so +// the ticker doesn't depend on any package-level state — production +// passes the daemon-owned Loader, tests pass a fake Loader (built with +// a fake PolicyFetcher). +// +// The initial snapshot is populated by calling loader.Load() at // construction time so the first tick only fires // onChange when the policy actually changed since boot — without // this baseline the first tick would report every currently-managed // key as "added" and trigger a spurious engine restart. -func NewTicker(reloadInterval time.Duration) *Ticker { +func NewTicker(reloadInterval time.Duration, loader *Loader) *Ticker { return &Ticker{ interval: reloadInterval, - prev: policyLoader(), + loader: loader, + prev: loader.Load(), } } @@ -58,13 +58,10 @@ func (t *Ticker) Run(ctx context.Context, onChange func(prev, curr *Policy) erro log.Info("MDM policy reload ticker stopped") return case <-tk.C: - curr := policyLoader() - if policiesEqual(t.prev, curr) { + curr := t.loader.Load() + if !policyChanged(t.prev, curr) { continue } - added, removed, changed := diffPolicies(t.prev, curr) - log.Infof("MDM policy changed: added=%v removed=%v changed=%v", - added, removed, changed) prev := t.prev if err := onChange(prev, curr); err != nil { log.Errorf("MDM policy change handler failed (retrying in 1 minute): %v", err) @@ -127,3 +124,12 @@ func mapOf(p *Policy) map[string]any { } return out } + +func policyChanged(prev, curr *Policy) bool { + if policiesEqual(prev, curr) { + return false + } + added, removed, changed := diffPolicies(prev, curr) + log.Infof("MDM policy changed: added=%v removed=%v changed=%v", added, removed, changed) + return true +} diff --git a/client/mdm/ticker_test.go b/client/mdm/ticker_test.go index 17f3cfc2f..29e48e728 100644 --- a/client/mdm/ticker_test.go +++ b/client/mdm/ticker_test.go @@ -13,28 +13,40 @@ import ( // testReloadInterval for speeding up the ticker cadence under `go test` const testReloadInterval = 1 * time.Second -// withPolicyLoader overrides the package-level policyLoader for the duration -// of the test so the ticker observes a scripted policy instead of the real -// OS-native store. The original loader is restored on cleanup. -func withPolicyLoader(t *testing.T, fn func() *Policy) { - t.Helper() - prev := policyLoader - policyLoader = fn - t.Cleanup(func() { policyLoader = prev }) +// fakePolicyFetcher implements PolicyFetcher returning a scripted +// policy map. Goroutine-safe so the test can mutate the script while +// the ticker is observing it. +type fakePolicyFetcher struct { + mu sync.Mutex + values map[string]any +} + +func (f *fakePolicyFetcher) Fetch() map[string]any { + f.mu.Lock() + defer f.mu.Unlock() + if f.values == nil { + return nil + } + out := make(map[string]any, len(f.values)) + for k, v := range f.values { + out[k] = v + } + return out +} + +func (f *fakePolicyFetcher) set(values map[string]any) { + f.mu.Lock() + defer f.mu.Unlock() + f.values = values } func TestTicker_FiresOnChangeWithDelta(t *testing.T) { - var mu sync.Mutex - current := NewPolicy(nil) // initial observation: empty (no enforcement) - withPolicyLoader(t, func() *Policy { - mu.Lock() - defer mu.Unlock() - return current - }) + fetcher := &fakePolicyFetcher{} // initial observation: empty (no enforcement) + loader := NewLoader(fetcher) type change struct{ prev, curr *Policy } changes := make(chan change, 1) - tk := NewTicker(testReloadInterval) + tk := NewTicker(testReloadInterval, loader) require.Equal(t, testReloadInterval, tk.interval) ctx, cancel := context.WithCancel(context.Background()) @@ -49,15 +61,13 @@ func TestTicker_FiresOnChangeWithDelta(t *testing.T) { }) close(done) }() - // Stop Run and wait for it to exit before returning, so the policyLoader - // restore in t.Cleanup can't race the ticker goroutine still reading it. + // Stop Run and wait for it to exit before returning, so the test + // goroutine doesn't race the still-running ticker. defer func() { cancel(); <-done }() - // Flip the OS-observed policy from empty to one managed key. The next - // tick must detect the diff and invoke onChange. - mu.Lock() - current = NewPolicy(map[string]any{KeyManagementURL: "https://mdm.example.com:443"}) - mu.Unlock() + // Flip the OS-observed policy from empty to one managed key. The + // next tick must detect the diff and invoke onChange. + fetcher.set(map[string]any{KeyManagementURL: "https://mdm.example.com:443"}) select { case c := <-changes: @@ -69,12 +79,11 @@ func TestTicker_FiresOnChangeWithDelta(t *testing.T) { } func TestTicker_NoCallbackWhenPolicyUnchanged(t *testing.T) { - withPolicyLoader(t, func() *Policy { - return NewPolicy(map[string]any{KeyBlockInbound: true}) - }) + fetcher := &fakePolicyFetcher{values: map[string]any{KeyBlockInbound: true}} + loader := NewLoader(fetcher) fired := make(chan struct{}, 1) - tk := NewTicker(testReloadInterval) + tk := NewTicker(testReloadInterval, loader) ctx, cancel := context.WithCancel(context.Background()) done := make(chan struct{}) @@ -90,8 +99,8 @@ func TestTicker_NoCallbackWhenPolicyUnchanged(t *testing.T) { }() defer func() { cancel(); <-done }() - // Over ~2 ticks at the 1s test cadence the policy never changes, so the - // diff guard must suppress the callback entirely. + // Over ~2 ticks at the 1s test cadence the policy never changes, + // so the diff guard must suppress the callback entirely. select { case <-fired: t.Fatal("onChange fired despite an unchanged policy") diff --git a/client/mobile/profile_lifecycle_test.go b/client/mobile/profile_lifecycle_test.go new file mode 100644 index 000000000..9612f550d --- /dev/null +++ b/client/mobile/profile_lifecycle_test.go @@ -0,0 +1,97 @@ +package mobile + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/internal/profilemanager" +) + +// loadAsTheMobileSDKsDo replays what the iOS SDK does with a stored profile: +// read the config, serialize it, and load it back. Client.SetConfigFromJSON +// stores that document for tvOS, Auth.SetConfigFromJSON authenticates with it, +// and copyConfig round-trips a Config through the same pair to take an +// in-memory copy before applying the MDM overlay. +func loadAsTheMobileSDKsDo(t *testing.T, configPath string) *profilemanager.Config { + t.Helper() + + stored, err := profilemanager.GetExistingConfig(configPath) + require.NoError(t, err, "read the stored profile") + + document, err := profilemanager.ConfigToJSON(stored) + require.NoError(t, err, "serialize the stored profile") + + reloaded, err := profilemanager.ConfigFromJSON(document) + require.NoError(t, err, "load the profile back") + return reloaded +} + +// A profile survives the whole round its user puts it through: created, logged +// out, loaded again, and switched away from and back. +// +// Logout is the step that makes this worth asserting. It clears the peer's +// keys in place so the next login registers a new peer rather than bringing +// the old one back, which leaves a profile that legitimately carries no +// identity — and both mobile SDKs go on loading that profile through the +// serialized form. A load that refused it, or a creation that never wrote an +// identity in the first place, breaks logout and profile switching on iOS and +// Android without any of it being visible from the desktop client. +func TestProfileSurvivesLogoutAndReload(t *testing.T) { + pm := newTestProfileManager(t) + + created, err := pm.AddProfile("work") + require.NoError(t, err) + require.NoError(t, pm.SwitchProfile(profilemanager.DefaultProfileName)) + + // Created: the profile carries the identity it will connect with. + require.NotEmpty(t, privateKeyOf(t, pm, created.ID), "a new profile was written with no identity") + + configPath, err := pm.GetConfigPath(created.ID) + require.NoError(t, err) + + before := loadAsTheMobileSDKsDo(t, configPath) + require.NotEmpty(t, before.PrivateKey) + managementURL := before.ManagementURL.String() + + // Logged out: the identity is gone, on purpose. + require.NoError(t, pm.LogoutProfile(created.ID)) + require.Empty(t, privateKeyOf(t, pm, created.ID), "logout left the peer's key behind") + + // Loaded again: the profile is still readable, and loading it neither + // fails nor mints a key that nothing would write down. + after := loadAsTheMobileSDKsDo(t, configPath) + assert.Empty(t, after.PrivateKey, "loading a logged-out profile minted a key nothing will persist") + assert.Empty(t, after.SSHKey, "loading a logged-out profile minted an SSH key") + assert.Equal(t, managementURL, after.ManagementURL.String(), "the rest of the profile did not survive the logout") + + // Switched away from and back: still the same profile, still loadable. + require.NoError(t, pm.SwitchProfile(created.ID)) + require.NoError(t, pm.SwitchProfile(profilemanager.DefaultProfileName)) + require.NoError(t, pm.SwitchProfile(created.ID)) + + active, err := pm.GetActiveProfile() + require.NoError(t, err) + assert.Equal(t, created.ID, active.ID, "the profile switched to is not the active one") + + assert.Equal(t, managementURL, loadAsTheMobileSDKsDo(t, configPath).ManagementURL.String(), + "the profile did not survive the round of switches") +} + +// The profile the SDKs fall back to gets the same treatment, since it is the +// one a mobile client without an explicit profile runs on. +func TestDefaultProfileSurvivesLogoutAndReload(t *testing.T) { + pm := newTestProfileManager(t) + require.NoError(t, pm.SwitchProfile(profilemanager.DefaultProfileName)) + + configPath, err := pm.GetConfigPath(profilemanager.DefaultProfileName) + require.NoError(t, err) + require.NotEmpty(t, loadAsTheMobileSDKsDo(t, configPath).PrivateKey) + + require.NoError(t, pm.LogoutProfile(profilemanager.DefaultProfileName)) + + reloaded := loadAsTheMobileSDKsDo(t, configPath) + assert.Empty(t, reloaded.PrivateKey, "loading the logged-out default profile minted a key") + assert.NotNil(t, reloaded.ManagementURL, "the profile lost its management URL") +} diff --git a/client/mobile/profile_manager.go b/client/mobile/profile_manager.go index 1ddabf0a9..ad79d80c0 100644 --- a/client/mobile/profile_manager.go +++ b/client/mobile/profile_manager.go @@ -4,6 +4,7 @@ package mobile import ( + "errors" "fmt" "os" "path/filepath" @@ -11,6 +12,7 @@ import ( log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/client/internal/profilemanager" + "github.com/netbirdio/netbird/client/mdm" ) const ( @@ -22,6 +24,9 @@ const ( profilesSubdir = "profiles" ) +// ErrProfilesDisabled marks a profile mutation rejected by MDM policy. +var ErrProfilesDisabled = errors.New("profile management is disabled by MDM policy") + /* / ← app-writable config root @@ -55,6 +60,7 @@ type ProfileManager struct { configDir string username string serviceMgr *profilemanager.ServiceManager + mdmLoader *mdm.Loader } // NewProfileManager creates a profile manager rooted at configDir, the @@ -127,6 +133,9 @@ func (pm *ProfileManager) GetActiveProfile() (*Profile, error) { // SwitchProfile records the given profile ID as the active profile. The caller // must stop the VPN tunnel before switching. func (pm *ProfileManager) SwitchProfile(id string) error { + if err := pm.checkProfilesAllowed(); err != nil { + return err + } if err := pm.serviceMgr.SetActiveProfileState(&profilemanager.ActiveProfileState{ ID: profilemanager.ID(id), Username: pm.username, @@ -141,6 +150,9 @@ func (pm *ProfileManager) SwitchProfile(id string) error { // AddProfile creates a new profile with the given display name and a // generated ID. It returns the created profile so the caller learns the ID. func (pm *ProfileManager) AddProfile(displayName string) (*Profile, error) { + if err := pm.checkProfilesAllowed(); err != nil { + return nil, err + } profile, err := pm.serviceMgr.AddProfile(displayName, pm.username) if err != nil { return nil, fmt.Errorf("add profile: %w", err) @@ -153,6 +165,9 @@ func (pm *ProfileManager) AddProfile(displayName string) (*Profile, error) { // RenameProfile changes the display name of the profile identified by id. The // on-disk filename (the ID) is left unchanged. func (pm *ProfileManager) RenameProfile(id string, newName string) error { + if err := pm.checkProfilesAllowed(); err != nil { + return err + } if err := pm.serviceMgr.RenameProfile(profilemanager.ID(id), pm.username, newName); err != nil { return fmt.Errorf("rename profile: %w", err) } @@ -165,6 +180,9 @@ func (pm *ProfileManager) RenameProfile(id string, newName string) error { // private key and SSH key from the config, forcing a re-login. The management // URL and other settings are preserved. func (pm *ProfileManager) LogoutProfile(id string) error { + if err := pm.checkProfileLogoutAllowed(id); err != nil { + return err + } configPath, err := pm.getProfileConfigPath(id) if err != nil { return err @@ -174,7 +192,10 @@ func (pm *ProfileManager) LogoutProfile(id string) error { return fmt.Errorf("profile %q does not exist", id) } - config, err := profilemanager.ReadConfig(configPath) + // The existing-file reader, not the generating one: the check above is not + // atomic with this read, so a profile removed in between would otherwise be + // resolved from the defaults here and recreated by the write below. + config, err := profilemanager.GetExistingConfig(configPath) if err != nil { return fmt.Errorf("read profile config: %w", err) } @@ -196,6 +217,9 @@ func (pm *ProfileManager) LogoutProfile(id string) error { // RemoveProfile deletes a profile. The default profile and the active profile // cannot be removed. func (pm *ProfileManager) RemoveProfile(id string) error { + if err := pm.checkProfilesAllowed(); err != nil { + return err + } configPath, err := pm.getProfileConfigPath(id) if err != nil { return err @@ -267,6 +291,27 @@ func (pm *ProfileManager) GetActiveStateFilePath() (string, error) { return pm.GetStateFilePath(activeProfile.ID) } +// SetMDMLoader registers the MDM policy source consulted before profile +// mutations; a nil loader disables enforcement. +func (pm *ProfileManager) SetMDMLoader(loader *mdm.Loader) { + pm.mdmLoader = loader +} + +func (pm *ProfileManager) checkProfilesAllowed() error { + if v, ok := pm.mdmLoader.Load().GetBool(mdm.KeyDisableProfiles); ok && v { + return ErrProfilesDisabled + } + return nil +} + +func (pm *ProfileManager) checkProfileLogoutAllowed(id string) error { + active, err := pm.serviceMgr.GetActiveProfileState() + if err == nil && active.ID.String() == id { + return nil + } + return pm.checkProfilesAllowed() +} + // profileEmail returns the account email recorded for a profile. Display-only, // so an unresolvable path degrades to "" rather than an error. func (pm *ProfileManager) profileEmail(id string) string { diff --git a/client/mobile/profile_manager_mdm_test.go b/client/mobile/profile_manager_mdm_test.go new file mode 100644 index 000000000..305becac3 --- /dev/null +++ b/client/mobile/profile_manager_mdm_test.go @@ -0,0 +1,83 @@ +package mobile + +import ( + "encoding/json" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/internal/profilemanager" + "github.com/netbirdio/netbird/client/mdm" +) + +type fakeFetcher struct{ values map[string]any } + +func (f *fakeFetcher) Fetch() map[string]any { return f.values } + +func newTestProfileManager(t *testing.T) *ProfileManager { + t.Helper() + origDir := profilemanager.DefaultConfigPathDir + origPath := profilemanager.DefaultConfigPath + origActive := profilemanager.ActiveProfileStatePath + t.Cleanup(func() { + profilemanager.DefaultConfigPathDir = origDir + profilemanager.DefaultConfigPath = origPath + profilemanager.ActiveProfileStatePath = origActive + }) + + configDir := t.TempDir() + pm := NewProfileManager(configDir, "mobile") + _, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{ + ConfigPath: filepath.Join(configDir, defaultConfigFilename), + }) + require.NoError(t, err) + return pm +} + +func privateKeyOf(t *testing.T, pm *ProfileManager, id string) string { + t.Helper() + path, err := pm.getProfileConfigPath(id) + require.NoError(t, err) + raw, err := os.ReadFile(path) + require.NoError(t, err) + var cfg struct{ PrivateKey string } + require.NoError(t, json.Unmarshal(raw, &cfg)) + return cfg.PrivateKey +} + +func TestLogoutProfile_DisableProfiles(t *testing.T) { + pm := newTestProfileManager(t) + other, err := pm.AddProfile("work") + require.NoError(t, err) + require.NoError(t, pm.SwitchProfile(profilemanager.DefaultProfileName)) + require.NotEmpty(t, privateKeyOf(t, pm, profilemanager.DefaultProfileName)) + require.NotEmpty(t, privateKeyOf(t, pm, other.ID)) + + pm.SetMDMLoader(mdm.NewLoader(&fakeFetcher{values: map[string]any{ + mdm.KeyDisableProfiles: true, + }})) + + err = pm.LogoutProfile(other.ID) + assert.ErrorIs(t, err, ErrProfilesDisabled) + assert.NotEmpty(t, privateKeyOf(t, pm, other.ID)) + + require.NoError(t, pm.LogoutProfile(profilemanager.DefaultProfileName)) + assert.Empty(t, privateKeyOf(t, pm, profilemanager.DefaultProfileName)) +} + +func TestLogoutProfile_ProfilesAllowed(t *testing.T) { + pm := newTestProfileManager(t) + other, err := pm.AddProfile("work") + require.NoError(t, err) + require.NoError(t, pm.SwitchProfile(profilemanager.DefaultProfileName)) + + pm.SetMDMLoader(mdm.NewLoader(&fakeFetcher{values: map[string]any{ + mdm.KeyDisableProfiles: false, + }})) + + require.NoError(t, pm.LogoutProfile(other.ID)) + assert.Empty(t, privateKeyOf(t, pm, other.ID)) +} diff --git a/client/netbird.wxs b/client/netbird.wxs index f30a7aa7e..156b4ff27 100644 --- a/client/netbird.wxs +++ b/client/netbird.wxs @@ -76,6 +76,14 @@ + + + +