mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-09 23:19:11 +02:00
Merge remote-tracking branch 'origin/main' into jnfrati/ubi-signal
This commit is contained in:
@@ -14,5 +14,15 @@ reviews:
|
|||||||
- "!**/*.ts"
|
- "!**/*.ts"
|
||||||
- "!**/*.js"
|
- "!**/*.js"
|
||||||
- "!**/*.svg"
|
- "!**/*.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:
|
chat:
|
||||||
auto_reply: true
|
auto_reply: true
|
||||||
|
|||||||
Executable
+26
@@ -0,0 +1,26 @@
|
|||||||
|
#!/bin/bash
|
||||||
|
# Refuses commit messages that carry attribution trailers. Contributors own
|
||||||
|
# their contributions (AGENTS.md, "No Co-Authored-By or tool-attribution
|
||||||
|
# trailers"); a trailer spreads that ownership onto a tool or a bystander.
|
||||||
|
|
||||||
|
msg_file="$1"
|
||||||
|
|
||||||
|
# Trailer keys in any casing, with any bullet or emoji in front.
|
||||||
|
trailers='^[^[:alnum:]]*(co-authored-by|claude-session|generated-by):'
|
||||||
|
# "Generated with/by" footers, including "Generated with <emoji> by".
|
||||||
|
footer='^[^[:alnum:]]*generated (with|by)( [^[:alnum:]]*by)? '
|
||||||
|
# A footer names a product, so a capitalized word must follow the phrase
|
||||||
|
# itself. Prose such as "generated by the protobuf compiler" stays legal.
|
||||||
|
tool='[Gg][Ee][Nn][Ee][Rr][Aa][Tt][Ee][Dd] ([Ww][Ii][Tt][Hh]|[Bb][Yy])( [^[:alnum:]]*[Bb][Yy])? [^[:alnum:]]*[A-Z]'
|
||||||
|
|
||||||
|
offending=$( {
|
||||||
|
grep -Ein "$trailers" "$msg_file"
|
||||||
|
grep -Ein "$footer" "$msg_file" | grep -E "$tool"
|
||||||
|
} | sort -un )
|
||||||
|
|
||||||
|
if [ -n "$offending" ]; then
|
||||||
|
echo "commit-msg: attribution trailers are not accepted in this repository:" >&2
|
||||||
|
printf '%s\n' "$offending" | sed 's/^/ /' >&2
|
||||||
|
echo "Remove them and commit again (see AGENTS.md)." >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
@@ -46,3 +46,25 @@ updates:
|
|||||||
wireguard:
|
wireguard:
|
||||||
patterns:
|
patterns:
|
||||||
- "golang.zx2c4.com/wireguard*"
|
- "golang.zx2c4.com/wireguard*"
|
||||||
|
|
||||||
|
# Base images of the source-build Dockerfiles, pinned by digest (Chainguard
|
||||||
|
# publishes only :latest for free). Dockerfile.release files feed goreleaser
|
||||||
|
# and keep the published images as they are, so their bases are left alone.
|
||||||
|
- package-ecosystem: "docker"
|
||||||
|
directories:
|
||||||
|
- "/upload-server"
|
||||||
|
schedule:
|
||||||
|
interval: "weekly"
|
||||||
|
open-pull-requests-limit: 3
|
||||||
|
groups:
|
||||||
|
base-images:
|
||||||
|
patterns:
|
||||||
|
- "*"
|
||||||
|
ignore:
|
||||||
|
- dependency-name: "gcr.io/distroless/base"
|
||||||
|
# Go minor and major versions move with the rest of the repository;
|
||||||
|
# patch releases and new digests of the pinned tag still come through.
|
||||||
|
- dependency-name: "golang"
|
||||||
|
update-types:
|
||||||
|
- "version-update:semver-minor"
|
||||||
|
- "version-update:semver-major"
|
||||||
|
|||||||
Executable
+338
@@ -0,0 +1,338 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
set -euo pipefail
|
||||||
|
|
||||||
|
fail() {
|
||||||
|
echo "::error::$*" >&2
|
||||||
|
exit 1
|
||||||
|
}
|
||||||
|
|
||||||
|
if [[ ${RUNNER_ENVIRONMENT:-} != github-hosted || ${RUNNER_OS:-} != macOS || $(uname -s) != Darwin ]]; then
|
||||||
|
fail "This test installs a system daemon and must run on a disposable GitHub macOS runner."
|
||||||
|
fi
|
||||||
|
if [[ $EUID == 0 ]]; then
|
||||||
|
fail "Run this script as the Homebrew user, not root."
|
||||||
|
fi
|
||||||
|
|
||||||
|
readonly test_dir="${RUNNER_TEMP:?}/homebrew-cask"
|
||||||
|
readonly results_dir="$test_dir/results"
|
||||||
|
readonly app='/Applications/Netbird UI.app'
|
||||||
|
readonly plist='/Library/LaunchDaemons/netbird.plist'
|
||||||
|
readonly cask='netbirdio/tap/netbird-ui'
|
||||||
|
readonly formula='netbirdio/tap/netbird'
|
||||||
|
readonly published_cask="$test_dir/published-netbird-ui.rb"
|
||||||
|
readonly legacy_cask="$test_dir/legacy-netbird-ui.rb"
|
||||||
|
readonly rendered_cask="$test_dir/rendered-netbird-ui.rb"
|
||||||
|
readonly fixture_dir="$test_dir/fixture"
|
||||||
|
readonly serve_dir="$test_dir/serve"
|
||||||
|
readonly fixture_zip="$serve_dir/netbird-ui.zip"
|
||||||
|
readonly fixture_port=18080
|
||||||
|
readonly fixture_url="http://127.0.0.1:$fixture_port/netbird-ui.zip"
|
||||||
|
readonly marker="$test_dir/installer.marker"
|
||||||
|
|
||||||
|
mkdir -p "$results_dir" "$fixture_dir/netbird_ui_darwin" "$serve_dir" "$test_dir/downloads"
|
||||||
|
exec > >(tee "$results_dir/test.log") 2>&1
|
||||||
|
|
||||||
|
sudo -n true
|
||||||
|
if command -v netbird || [[ -e "$app" || -e "$plist" ]] || pgrep -x netbird-ui; then
|
||||||
|
fail "The runner already has NetBird installed or running."
|
||||||
|
fi
|
||||||
|
if sudo launchctl print system/netbird > "$results_dir/initial-service.log" 2>&1; then
|
||||||
|
fail "The runner already has a NetBird service loaded."
|
||||||
|
fi
|
||||||
|
|
||||||
|
install_attempted=false
|
||||||
|
server_pid=''
|
||||||
|
daemon_pid=''
|
||||||
|
version=''
|
||||||
|
|
||||||
|
stop_ui() {
|
||||||
|
local status=0
|
||||||
|
sudo pkill -x netbird-ui || status=$?
|
||||||
|
# pkill returns 1 when the UI is already closed.
|
||||||
|
[[ $status == 0 || $status == 1 ]]
|
||||||
|
}
|
||||||
|
|
||||||
|
cleanup() {
|
||||||
|
local status=$?
|
||||||
|
trap - EXIT
|
||||||
|
set +e
|
||||||
|
|
||||||
|
if [[ $install_attempted == true ]]; then
|
||||||
|
stop_ui || status=1
|
||||||
|
if [[ -S /var/run/netbird.sock ]]; then
|
||||||
|
sudo netbird down || status=1
|
||||||
|
fi
|
||||||
|
if brew list --cask "$cask" >/dev/null 2>&1 || [[ -e "$app" ]]; then
|
||||||
|
brew uninstall --cask --force "$cask" || status=1
|
||||||
|
fi
|
||||||
|
# A failed cask install can leave a daemon even after Homebrew rolls back the app.
|
||||||
|
if sudo launchctl print system/netbird > "$results_dir/cleanup-service.log" 2>&1; then
|
||||||
|
sudo netbird service stop || status=1
|
||||||
|
fi
|
||||||
|
if [[ -e "$plist" ]]; then
|
||||||
|
sudo netbird service uninstall || status=1
|
||||||
|
fi
|
||||||
|
fi
|
||||||
|
if [[ -f /var/log/netbird/client.log ]]; then
|
||||||
|
sudo cat /var/log/netbird/client.log > "$results_dir/client.log" || status=1
|
||||||
|
fi
|
||||||
|
if command -v netbird >/dev/null; then
|
||||||
|
brew uninstall --formula "$formula" || status=1
|
||||||
|
fi
|
||||||
|
if [[ -n $server_pid ]]; then
|
||||||
|
kill "$server_pid" 2>/dev/null || true
|
||||||
|
fi
|
||||||
|
exit "$status"
|
||||||
|
}
|
||||||
|
trap cleanup EXIT
|
||||||
|
trap 'exit 130' INT
|
||||||
|
trap 'exit 143' TERM
|
||||||
|
|
||||||
|
run_logged() {
|
||||||
|
local name=$1
|
||||||
|
shift
|
||||||
|
"$@" 2>&1 | tee "$results_dir/$name.log"
|
||||||
|
}
|
||||||
|
|
||||||
|
cask_field() {
|
||||||
|
local stanza=$1 file=$2
|
||||||
|
sed -nE "s/^[[:space:]]*$stanza \"([^\"]+)\".*/\\1/p" "$file"
|
||||||
|
}
|
||||||
|
|
||||||
|
release_fields() {
|
||||||
|
local file=$1
|
||||||
|
grep -E '^[[:space:]]*(version|url|sha256|app) ' "$file"
|
||||||
|
}
|
||||||
|
|
||||||
|
use_cask() {
|
||||||
|
local file=$1
|
||||||
|
cp "$file" "$tap_dir/Casks/netbird-ui.rb"
|
||||||
|
}
|
||||||
|
|
||||||
|
# The released installer opens the UI as root, which never returns on a headless
|
||||||
|
# runner. The cask only needs two script paths and a version argument, so the test
|
||||||
|
# ships a stub bundle that records what it received and starts the daemon.
|
||||||
|
build_fixture() {
|
||||||
|
local bundle="$fixture_dir/netbird_ui_darwin"
|
||||||
|
printf '#!/bin/sh\nexit 0\n' > "$bundle/netbird-ui"
|
||||||
|
chmod 755 "$bundle/netbird-ui"
|
||||||
|
# After a bootout launchd keeps tearing the previous daemon down for a couple of
|
||||||
|
# seconds, and loading the same label again fails until that finishes.
|
||||||
|
cat > "$bundle/installer.sh" <<EOF
|
||||||
|
#!/bin/sh
|
||||||
|
set -eu
|
||||||
|
export PATH=\$PATH:/usr/local/bin:/opt/homebrew/bin
|
||||||
|
printf 'version=%s\\nuid=%s\\n' "\$1" "\$(id -u)" > '$marker'
|
||||||
|
netbird service install
|
||||||
|
attempt=0
|
||||||
|
until netbird service start; do
|
||||||
|
attempt=\$((attempt + 1))
|
||||||
|
[ "\$attempt" -lt 15 ] || exit 1
|
||||||
|
sleep 1
|
||||||
|
done
|
||||||
|
EOF
|
||||||
|
printf '#!/bin/sh\nexit 0\n' > "$bundle/uninstaller.sh"
|
||||||
|
# Shipped without the executable bit so the 0755 seen after install can only come from the cask.
|
||||||
|
chmod 644 "$bundle/installer.sh" "$bundle/uninstaller.sh"
|
||||||
|
rm -f "$fixture_zip"
|
||||||
|
(cd "$fixture_dir" && zip -qr "$fixture_zip" netbird_ui_darwin)
|
||||||
|
}
|
||||||
|
|
||||||
|
start_fixture_server() {
|
||||||
|
python3 -m http.server "$fixture_port" --bind 127.0.0.1 --directory "$serve_dir" \
|
||||||
|
> "$results_dir/fixture-server.log" 2>&1 &
|
||||||
|
server_pid=$!
|
||||||
|
local attempt
|
||||||
|
for attempt in {1..20}; do
|
||||||
|
if curl --silent --fail --output /dev/null "$fixture_url"; then
|
||||||
|
return
|
||||||
|
fi
|
||||||
|
sleep 0.5
|
||||||
|
done
|
||||||
|
fail "The fixture HTTP server did not come up on port $fixture_port."
|
||||||
|
}
|
||||||
|
|
||||||
|
assert_published_layout() {
|
||||||
|
local url archive script
|
||||||
|
while read -r url; do
|
||||||
|
archive="$test_dir/downloads/${url##*/}"
|
||||||
|
curl --fail --location --silent --retry 3 --output "$archive" "$url"
|
||||||
|
for script in installer.sh uninstaller.sh; do
|
||||||
|
unzip -l "$archive" | grep -q " netbird_ui_darwin/$script\$" ||
|
||||||
|
fail "The published archive ${url##*/} has no netbird_ui_darwin/$script."
|
||||||
|
done
|
||||||
|
done < <(cask_field url "$published_cask")
|
||||||
|
}
|
||||||
|
|
||||||
|
assert_no_deprecations() {
|
||||||
|
if grep -Ei '(postflight|uninstall_preflight).*deprecated|deprecated.*(postflight|uninstall_preflight)' "$@"; then
|
||||||
|
fail "Homebrew reported a deprecated cask lifecycle hook."
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
wait_for_daemon() {
|
||||||
|
local attempt
|
||||||
|
for attempt in {1..30}; do
|
||||||
|
if sudo launchctl print system/netbird > "$results_dir/service.log" 2>&1 &&
|
||||||
|
grep -Eq '^[[:space:]]*state = running$' "$results_dir/service.log"; then
|
||||||
|
return
|
||||||
|
fi
|
||||||
|
sleep 1
|
||||||
|
done
|
||||||
|
cat "$results_dir/service.log"
|
||||||
|
fail "The installed daemon did not reach the running state."
|
||||||
|
}
|
||||||
|
|
||||||
|
wait_for_exit() {
|
||||||
|
local pid=$1 attempt
|
||||||
|
for attempt in {1..30}; do
|
||||||
|
if ! sudo kill -0 "$pid" 2>/dev/null; then
|
||||||
|
return
|
||||||
|
fi
|
||||||
|
sleep 1
|
||||||
|
done
|
||||||
|
fail "Daemon process $pid is still running after removal."
|
||||||
|
}
|
||||||
|
|
||||||
|
assert_service_absent() {
|
||||||
|
if sudo launchctl print system/netbird > "$results_dir/removed-service.log" 2>&1; then
|
||||||
|
fail "The NetBird service is still loaded after removal."
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
assert_installed() {
|
||||||
|
local script
|
||||||
|
[[ -f $marker ]] || fail "The cask did not run installer.sh."
|
||||||
|
grep -qx "version=$version" "$marker" || fail "installer.sh did not receive the cask version: $(cat "$marker")"
|
||||||
|
grep -qx 'uid=0' "$marker" || fail "installer.sh did not run as root: $(cat "$marker")"
|
||||||
|
[[ -d "$app" && -x "$app/netbird-ui" ]] || fail "The UI was not installed."
|
||||||
|
for script in installer.sh uninstaller.sh; do
|
||||||
|
[[ $(stat -f '%Lp' "$app/$script") == 755 ]] || fail "Incorrect permissions on $script."
|
||||||
|
done
|
||||||
|
[[ -f "$plist" ]] || fail "The installer did not create the daemon plist."
|
||||||
|
wait_for_daemon
|
||||||
|
daemon_pid=$(awk '/^[[:space:]]*pid = / { print $3; exit }' "$results_dir/service.log")
|
||||||
|
[[ $daemon_pid =~ ^[0-9]+$ ]] || fail "The running daemon has no PID."
|
||||||
|
sudo kill -0 "$daemon_pid"
|
||||||
|
}
|
||||||
|
|
||||||
|
assert_uninstalled() {
|
||||||
|
local log=$1
|
||||||
|
assert_no_deprecations "$log"
|
||||||
|
[[ ! -e "$app" ]] || fail "The UI app remains after uninstall."
|
||||||
|
[[ ! -e "$plist" ]] || fail "The daemon plist remains after uninstall."
|
||||||
|
assert_service_absent
|
||||||
|
wait_for_exit "$daemon_pid"
|
||||||
|
[[ $(netbird version) == "$version" ]] || fail "Cask uninstall removed the CLI dependency."
|
||||||
|
}
|
||||||
|
|
||||||
|
installed_caskfiles() {
|
||||||
|
local extension=$1
|
||||||
|
find "$(brew --caskroom)/netbird-ui/.metadata" -name "netbird-ui.$extension" 2>/dev/null
|
||||||
|
}
|
||||||
|
|
||||||
|
assert_legacy_metadata() {
|
||||||
|
installed_caskfiles rb | grep -q . || fail "The legacy cask did not leave a Ruby caskfile behind."
|
||||||
|
}
|
||||||
|
|
||||||
|
assert_steps_metadata() {
|
||||||
|
if installed_caskfiles rb | grep -q .; then
|
||||||
|
fail "Homebrew still keeps the legacy Ruby caskfile after reinstall."
|
||||||
|
fi
|
||||||
|
installed_caskfiles json | grep -q . || fail "Homebrew did not save the reinstalled cask as JSON."
|
||||||
|
}
|
||||||
|
|
||||||
|
brew --version
|
||||||
|
sw_vers
|
||||||
|
brew tap netbirdio/tap "${GITHUB_WORKSPACE:?}/.homebrew-cask-tap"
|
||||||
|
tap_dir=$(brew --repository netbirdio/tap)
|
||||||
|
readonly tap_dir
|
||||||
|
|
||||||
|
[[ -f "$tap_dir/Casks/netbird-ui.rb" ]] || fail "The tap has no Casks/netbird-ui.rb."
|
||||||
|
cp "$tap_dir/Casks/netbird-ui.rb" "$published_cask"
|
||||||
|
cp "$published_cask" "$results_dir/published-netbird-ui.rb"
|
||||||
|
|
||||||
|
version=$(brew info --json=v2 --formula "$formula" | jq -r '.formulae[0].versions.stable')
|
||||||
|
readonly version
|
||||||
|
[[ -n $version && $version != null ]] || fail "Could not read the formula version from the tap."
|
||||||
|
|
||||||
|
assert_published_layout
|
||||||
|
|
||||||
|
build_fixture
|
||||||
|
fixture_sha=$(shasum -a 256 "$fixture_zip" | cut -d' ' -f1)
|
||||||
|
readonly fixture_sha
|
||||||
|
start_fixture_server
|
||||||
|
|
||||||
|
export PROJECT=netbird-ui VERSION="$version"
|
||||||
|
export AMD="$fixture_zip" ARM="$fixture_zip" AMD_URL="$fixture_url" ARM_URL="$fixture_url"
|
||||||
|
gomplate -f "$GITHUB_WORKSPACE/client/ui/netbird-ui.rb.tmpl" -o "$rendered_cask"
|
||||||
|
cp "$rendered_cask" "$results_dir/rendered-netbird-ui.rb"
|
||||||
|
|
||||||
|
sed -E "s|^([[:space:]]*version) \"[^\"]+\"|\\1 \"$version\"|; s|^([[:space:]]*url) \"[^\"]+\"|\\1 \"$fixture_url\"|; s|^([[:space:]]*sha256) \"[^\"]+\"|\\1 \"$fixture_sha\"|" \
|
||||||
|
"$published_cask" > "$legacy_cask"
|
||||||
|
cp "$legacy_cask" "$results_dir/legacy-netbird-ui.rb"
|
||||||
|
if ! diff <(release_fields "$legacy_cask") <(release_fields "$rendered_cask"); then
|
||||||
|
fail "The rendered cask changes release data, not only lifecycle stanzas."
|
||||||
|
fi
|
||||||
|
|
||||||
|
use_cask "$rendered_cask"
|
||||||
|
brew info --json=v2 --cask "$cask" > "$results_dir/cask.json" 2> "$results_dir/load.log"
|
||||||
|
cat "$results_dir/load.log"
|
||||||
|
assert_no_deprecations "$results_dir/load.log"
|
||||||
|
run_logged style brew style --cask --only-cops=Cask/InstallSteps "$cask"
|
||||||
|
|
||||||
|
run_logged install-cli brew install --formula "$formula"
|
||||||
|
[[ $(netbird version) == "$version" ]] || fail "The installed CLI does not report the formula version."
|
||||||
|
|
||||||
|
for scenario in running stopped missing; do
|
||||||
|
echo "::group::Uninstall with $scenario service"
|
||||||
|
install_attempted=true
|
||||||
|
sudo rm -f "$marker"
|
||||||
|
run_logged "install-$scenario" brew install --cask "$cask"
|
||||||
|
assert_no_deprecations "$results_dir/install-$scenario.log"
|
||||||
|
assert_installed
|
||||||
|
stop_ui
|
||||||
|
|
||||||
|
case "$scenario" in
|
||||||
|
running) ;;
|
||||||
|
stopped)
|
||||||
|
run_logged stop-daemon sudo netbird service stop
|
||||||
|
wait_for_exit "$daemon_pid"
|
||||||
|
[[ -f "$plist" ]] || fail "Stopping the daemon unexpectedly removed its plist."
|
||||||
|
;;
|
||||||
|
missing)
|
||||||
|
run_logged stop-missing-daemon sudo netbird service stop
|
||||||
|
run_logged remove-daemon sudo netbird service uninstall
|
||||||
|
wait_for_exit "$daemon_pid"
|
||||||
|
[[ ! -e "$plist" ]] || fail "The missing-service scenario still has a plist."
|
||||||
|
assert_service_absent
|
||||||
|
;;
|
||||||
|
*) fail "Unknown uninstall scenario: $scenario" ;;
|
||||||
|
esac
|
||||||
|
|
||||||
|
run_logged "uninstall-$scenario" brew uninstall --cask "$cask"
|
||||||
|
assert_uninstalled "$results_dir/uninstall-$scenario.log"
|
||||||
|
echo "::endgroup::"
|
||||||
|
done
|
||||||
|
|
||||||
|
# Every existing user first meets the new cask through an upgrade of the published
|
||||||
|
# one, whose legacy flight blocks Homebrew replays from the saved Ruby caskfile.
|
||||||
|
echo "::group::Reinstall over the published legacy cask"
|
||||||
|
install_attempted=true
|
||||||
|
use_cask "$legacy_cask"
|
||||||
|
sudo rm -f "$marker"
|
||||||
|
run_logged install-legacy brew install --cask "$cask"
|
||||||
|
assert_installed
|
||||||
|
assert_legacy_metadata
|
||||||
|
stop_ui
|
||||||
|
|
||||||
|
use_cask "$rendered_cask"
|
||||||
|
sudo rm -f "$marker"
|
||||||
|
run_logged reinstall-legacy brew reinstall --cask "$cask"
|
||||||
|
assert_installed
|
||||||
|
assert_steps_metadata
|
||||||
|
stop_ui
|
||||||
|
|
||||||
|
run_logged uninstall-legacy brew uninstall --cask "$cask"
|
||||||
|
assert_uninstalled "$results_dir/uninstall-legacy.log"
|
||||||
|
echo "::endgroup::"
|
||||||
@@ -34,7 +34,7 @@ jobs:
|
|||||||
while IFS= read -r dir; do
|
while IFS= read -r dir; do
|
||||||
echo "=== Checking $dir ==="
|
echo "=== Checking $dir ==="
|
||||||
# Search for problematic imports, excluding test files
|
# 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
|
if [ -n "$RESULTS" ]; then
|
||||||
echo "❌ Found problematic dependencies:"
|
echo "❌ Found problematic dependencies:"
|
||||||
echo "$RESULTS"
|
echo "$RESULTS"
|
||||||
@@ -93,7 +93,7 @@ jobs:
|
|||||||
IMPORTERS=$(go list -json -deps ./... 2>/dev/null | jq -r "select(.Imports[]? == \"$package\") | .ImportPath")
|
IMPORTERS=$(go list -json -deps ./... 2>/dev/null | jq -r "select(.Imports[]? == \"$package\") | .ImportPath")
|
||||||
|
|
||||||
# Check if any importer is NOT in management/signal/relay
|
# 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
|
if [ -n "$BSD_IMPORTER" ]; then
|
||||||
echo "❌ $package ($license) is imported by BSD-licensed code: $BSD_IMPORTER"
|
echo "❌ $package ($license) is imported by BSD-licensed code: $BSD_IMPORTER"
|
||||||
|
|||||||
@@ -12,6 +12,8 @@ jobs:
|
|||||||
docs-ack:
|
docs-ack:
|
||||||
name: Require docs PR URL or explicit "not needed"
|
name: Require docs PR URL or explicit "not needed"
|
||||||
runs-on: ubuntu-latest
|
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:
|
steps:
|
||||||
- name: Read PR body
|
- name: Read PR body
|
||||||
|
|||||||
@@ -38,12 +38,12 @@ jobs:
|
|||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
- name: Set up Node.js
|
- name: Set up Node.js
|
||||||
uses: actions/setup-node@v4
|
uses: actions/setup-node@v7
|
||||||
with:
|
with:
|
||||||
node-version: "22"
|
node-version: "22"
|
||||||
|
|
||||||
- name: Set up pnpm
|
- name: Set up pnpm
|
||||||
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
|
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
|
||||||
with:
|
with:
|
||||||
version: 11
|
version: 11
|
||||||
|
|
||||||
@@ -79,7 +79,7 @@ jobs:
|
|||||||
run: echo "path=$(pnpm store path --silent)" >> "$GITHUB_OUTPUT"
|
run: echo "path=$(pnpm store path --silent)" >> "$GITHUB_OUTPUT"
|
||||||
|
|
||||||
- name: Cache pnpm store
|
- name: Cache pnpm store
|
||||||
uses: actions/cache@v4
|
uses: actions/cache@v6
|
||||||
with:
|
with:
|
||||||
path: ${{ steps.pnpm-store.outputs.path }}
|
path: ${{ steps.pnpm-store.outputs.path }}
|
||||||
key: ${{ runner.os }}-pnpm-${{ hashFiles('client/ui/frontend/pnpm-lock.yaml') }}
|
key: ${{ runner.os }}-pnpm-${{ hashFiles('client/ui/frontend/pnpm-lock.yaml') }}
|
||||||
|
|||||||
@@ -46,15 +46,17 @@ jobs:
|
|||||||
run: git --no-pager diff --exit-code
|
run: git --no-pager diff --exit-code
|
||||||
|
|
||||||
- name: Test
|
- name: Test
|
||||||
# Exclude client/ui: its main.go uses //go:embed all:frontend/dist,
|
# Exclude the client/ui package itself: its main.go uses //go:embed
|
||||||
# which fails to compile until the frontend has been built. The Wails UI
|
# all:frontend/dist, which fails to compile until the frontend has been
|
||||||
# has no Go-side unit tests, and its release pipeline runs `pnpm build`
|
# built, and its release pipeline runs `pnpm build` before goreleaser.
|
||||||
# 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
|
# `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,
|
# 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
|
# go list aborts with empty stdout and `go test` falls back to the repo
|
||||||
# root, which has no Go files.
|
# 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
|
- name: Upload coverage reports to Codecov
|
||||||
uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f #v7.0.0
|
uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f #v7.0.0
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
usesh: true
|
usesh: true
|
||||||
copyback: false
|
copyback: false
|
||||||
release: "15.0"
|
release: "15.1"
|
||||||
envs: "GO_VERSION"
|
envs: "GO_VERSION"
|
||||||
prepare: |
|
prepare: |
|
||||||
pkg install -y curl pkgconf xorg
|
pkg install -y curl pkgconf xorg
|
||||||
|
|||||||
@@ -160,9 +160,10 @@ jobs:
|
|||||||
|
|
||||||
- name: Test
|
- name: Test
|
||||||
# Exclude client/ui: its main.go uses //go:embed all:frontend/dist,
|
# 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
|
# which fails to compile until the frontend has been built, and its
|
||||||
# has no Go-side unit tests, and its release pipeline runs `pnpm build`
|
# release pipeline runs `pnpm build` before goreleaser. The subpackages
|
||||||
# before goreleaser.
|
# 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
|
# `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,
|
# 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
|
# go list aborts with empty stdout and `go test` falls back to the repo
|
||||||
@@ -177,6 +178,35 @@ jobs:
|
|||||||
slug: netbirdio/netbird
|
slug: netbirdio/netbird
|
||||||
flags: unit,client
|
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:
|
test_client_on_docker:
|
||||||
name: "Client (Docker) / Unit"
|
name: "Client (Docker) / Unit"
|
||||||
needs: [build-cache]
|
needs: [build-cache]
|
||||||
@@ -211,6 +241,9 @@ jobs:
|
|||||||
${{ runner.os }}-gotest-cache-
|
${{ runner.os }}-gotest-cache-
|
||||||
|
|
||||||
- name: Run tests in container
|
- 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:
|
env:
|
||||||
HOST_GOCACHE: ${{ steps.go-env.outputs.cache_dir }}
|
HOST_GOCACHE: ${{ steps.go-env.outputs.cache_dir }}
|
||||||
HOST_GOMODCACHE: ${{ steps.go-env.outputs.modcache_dir }}
|
HOST_GOMODCACHE: ${{ steps.go-env.outputs.modcache_dir }}
|
||||||
@@ -237,7 +270,7 @@ jobs:
|
|||||||
sh -c ' \
|
sh -c ' \
|
||||||
apk update; apk add --no-cache \
|
apk update; apk add --no-cache \
|
||||||
ca-certificates iptables ip6tables dbus dbus-dev libpcap-dev build-base; \
|
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:
|
test_relay:
|
||||||
@@ -481,14 +514,32 @@ jobs:
|
|||||||
if: matrix.store == 'mysql'
|
if: matrix.store == 'mysql'
|
||||||
run: docker pull mlsmaycon/warmed-mysql:8
|
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
|
- name: Test
|
||||||
|
shell: bash
|
||||||
run: |
|
run: |
|
||||||
|
set -o pipefail
|
||||||
CGO_ENABLED=1 GOARCH=${{ matrix.arch }} \
|
CGO_ENABLED=1 GOARCH=${{ matrix.arch }} \
|
||||||
NETBIRD_STORE_ENGINE=${{ matrix.store }} \
|
NETBIRD_STORE_ENGINE=${{ matrix.store }} \
|
||||||
CI=true \
|
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" \
|
-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
|
- name: Upload coverage reports to Codecov
|
||||||
if: matrix.arch == 'amd64'
|
if: matrix.arch == 'amd64'
|
||||||
@@ -738,12 +789,27 @@ jobs:
|
|||||||
- name: check git status
|
- name: check git status
|
||||||
run: git --no-pager diff --exit-code
|
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
|
- name: Test
|
||||||
|
shell: bash
|
||||||
run: |
|
run: |
|
||||||
|
set -o pipefail
|
||||||
CGO_ENABLED=1 GOARCH=${{ matrix.arch }} \
|
CGO_ENABLED=1 GOARCH=${{ matrix.arch }} \
|
||||||
NETBIRD_STORE_ENGINE=${{ matrix.store }} \
|
NETBIRD_STORE_ENGINE=${{ matrix.store }} \
|
||||||
CI=true \
|
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
|
- name: Upload coverage reports to Codecov
|
||||||
if: matrix.arch == 'amd64'
|
if: matrix.arch == 'amd64'
|
||||||
|
|||||||
@@ -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 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
|
- 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
|
- name: Generate test script
|
||||||
# Exclude client/ui: its main.go uses //go:embed all:frontend/dist,
|
# Exclude the client/ui package itself: its main.go uses //go:embed
|
||||||
# which fails to compile until the frontend has been built. The Wails UI
|
# all:frontend/dist, which fails to compile until the frontend has been
|
||||||
# has no Go-side unit tests, and its release pipeline runs `pnpm build`
|
# built, and its release pipeline runs `pnpm build` before goreleaser.
|
||||||
# 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
|
# `go list -e` lets the listing succeed even though the embed fails to
|
||||||
# resolve; the Where-Object pipeline then drops the broken package by
|
# resolve; the Where-Object pipeline then drops the broken package by
|
||||||
# path. Without -e, go list aborts with empty stdout.
|
# path. Without -e, go list aborts with empty stdout.
|
||||||
run: |
|
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"
|
$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"
|
$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
|
Set-Content -Path "${{ github.workspace }}\run-tests.cmd" -Value $cmd
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ jobs:
|
|||||||
# segment by codespell and behave the same across versions; the
|
# segment by codespell and behave the same across versions; the
|
||||||
# recursive "**" form did not take effect with the codespell shipped
|
# recursive "**" form did not take effect with the codespell shipped
|
||||||
# by this action.
|
# 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:
|
golangci:
|
||||||
strategy:
|
strategy:
|
||||||
fail-fast: false
|
fail-fast: false
|
||||||
@@ -80,3 +80,49 @@ jobs:
|
|||||||
skip-save-cache: true
|
skip-save-cache: true
|
||||||
cache-invalidation-interval: 0
|
cache-invalidation-interval: 0
|
||||||
args: --timeout=20m
|
args: --timeout=20m
|
||||||
|
|
||||||
|
# Separate job rather than extra rows in the matrix above: those rows pick a
|
||||||
|
# GOOS by picking a runner OS, while android/ios are cross-compiled from
|
||||||
|
# ubuntu — an `include` entry with os: ubuntu-latest would merge into the
|
||||||
|
# Linux row instead of adding one. The package path is restricted because a
|
||||||
|
# whole-repo run under GOOS=android pulls *_linux.go files into packages that
|
||||||
|
# have no android counterpart.
|
||||||
|
golangci-mobile:
|
||||||
|
strategy:
|
||||||
|
fail-fast: false
|
||||||
|
matrix:
|
||||||
|
include:
|
||||||
|
- goos: android
|
||||||
|
goarch: arm64
|
||||||
|
packages: ./client/android/...
|
||||||
|
display_name: Android
|
||||||
|
- goos: ios
|
||||||
|
goarch: arm64
|
||||||
|
packages: ./client/ios/...
|
||||||
|
display_name: iOS
|
||||||
|
name: ${{ matrix.display_name }}
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
timeout-minutes: 25
|
||||||
|
env:
|
||||||
|
CGO_ENABLED: 0
|
||||||
|
GOOS: ${{ matrix.goos }}
|
||||||
|
GOARCH: ${{ matrix.goarch }}
|
||||||
|
steps:
|
||||||
|
- name: Checkout code
|
||||||
|
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
|
||||||
|
with:
|
||||||
|
persist-credentials: false
|
||||||
|
- name: Install Go
|
||||||
|
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
|
||||||
|
with:
|
||||||
|
go-version-file: "go.mod"
|
||||||
|
cache: false
|
||||||
|
- name: golangci-lint
|
||||||
|
uses: golangci/golangci-lint-action@82606bf257cbaff209d206a39f5134f0cfbfd2ee #v9.2.1
|
||||||
|
with:
|
||||||
|
version: latest
|
||||||
|
install-mode: binary
|
||||||
|
skip-cache: true
|
||||||
|
skip-save-cache: true
|
||||||
|
cache-invalidation-interval: 0
|
||||||
|
args: --timeout=20m ${{ matrix.packages }}
|
||||||
|
|||||||
@@ -0,0 +1,64 @@
|
|||||||
|
name: Mobile
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
branches:
|
||||||
|
- main
|
||||||
|
- "release-*"
|
||||||
|
pull_request:
|
||||||
|
|
||||||
|
concurrency:
|
||||||
|
group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }}
|
||||||
|
cancel-in-progress: true
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
android_build:
|
||||||
|
name: "Android / Build"
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
strategy:
|
||||||
|
fail-fast: false
|
||||||
|
matrix:
|
||||||
|
goarch: [arm64, arm, amd64, "386"]
|
||||||
|
env:
|
||||||
|
CGO_ENABLED: 0
|
||||||
|
GOOS: android
|
||||||
|
GOARCH: ${{ matrix.goarch }}
|
||||||
|
steps:
|
||||||
|
- name: Checkout repository
|
||||||
|
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
|
||||||
|
with:
|
||||||
|
persist-credentials: false
|
||||||
|
- name: Install Go
|
||||||
|
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
|
||||||
|
with:
|
||||||
|
go-version-file: "go.mod"
|
||||||
|
- name: Build Android bridge
|
||||||
|
run: go build ./client/android/...
|
||||||
|
- name: Vet Android bridge
|
||||||
|
if: matrix.goarch == 'arm64'
|
||||||
|
run: go vet ./client/android/...
|
||||||
|
|
||||||
|
ios_build:
|
||||||
|
name: "iOS / Build"
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
strategy:
|
||||||
|
fail-fast: false
|
||||||
|
matrix:
|
||||||
|
goarch: [arm64, amd64]
|
||||||
|
env:
|
||||||
|
CGO_ENABLED: 0
|
||||||
|
GOOS: ios
|
||||||
|
GOARCH: ${{ matrix.goarch }}
|
||||||
|
steps:
|
||||||
|
- name: Checkout repository
|
||||||
|
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
|
||||||
|
with:
|
||||||
|
persist-credentials: false
|
||||||
|
- name: Install Go
|
||||||
|
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
|
||||||
|
with:
|
||||||
|
go-version-file: "go.mod"
|
||||||
|
# No `go vet` counterpart: every ios target requires external (cgo)
|
||||||
|
# linking, which needs an Xcode toolchain the runner does not have.
|
||||||
|
- name: Build iOS SDK
|
||||||
|
run: go build ./client/ios/...
|
||||||
@@ -7,6 +7,8 @@ on:
|
|||||||
jobs:
|
jobs:
|
||||||
check-title:
|
check-title:
|
||||||
runs-on: ubuntu-latest
|
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:
|
steps:
|
||||||
- name: Validate PR title prefix
|
- name: Validate PR title prefix
|
||||||
uses: actions/github-script@3a2844b7e9c422d3c10d287c895573f7108da1b3 # v9.0.0
|
uses: actions/github-script@3a2844b7e9c422d3c10d287c895573f7108da1b3 # v9.0.0
|
||||||
|
|||||||
@@ -0,0 +1,201 @@
|
|||||||
|
name: Red Hat Certification
|
||||||
|
|
||||||
|
# Certify published UBI images in the Red Hat Ecosystem Catalog. Called by
|
||||||
|
# release.yml on stable tags, or run by hand to (re)certify any released
|
||||||
|
# version. preflight submits every architecture of an image's manifest list
|
||||||
|
# to Pyxis; auto-publish on the component makes it public once certified.
|
||||||
|
#
|
||||||
|
# Each component's Partner Connect ID comes from the REDHAT_CERT_ID_<NAME>
|
||||||
|
# repository variable, e.g. REDHAT_CERT_ID_CLIENT_ROOTLESS. The run fails
|
||||||
|
# before certifying anything if a selected component's variable is not set.
|
||||||
|
|
||||||
|
on:
|
||||||
|
workflow_call:
|
||||||
|
inputs:
|
||||||
|
component:
|
||||||
|
type: string
|
||||||
|
required: true
|
||||||
|
version:
|
||||||
|
type: string
|
||||||
|
required: true
|
||||||
|
secrets:
|
||||||
|
PYXIS_API_TOKEN:
|
||||||
|
required: true
|
||||||
|
workflow_dispatch:
|
||||||
|
inputs:
|
||||||
|
component:
|
||||||
|
description: "Component to certify"
|
||||||
|
type: choice
|
||||||
|
required: true
|
||||||
|
default: all
|
||||||
|
options:
|
||||||
|
- all
|
||||||
|
- client-rootless
|
||||||
|
- reverse-proxy
|
||||||
|
- netbird-server
|
||||||
|
version:
|
||||||
|
description: "Released version, e.g. v0.80.0"
|
||||||
|
type: string
|
||||||
|
required: true
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
resolve:
|
||||||
|
name: Resolve components
|
||||||
|
runs-on: ubuntu-24.04
|
||||||
|
outputs:
|
||||||
|
version: ${{ steps.resolve.outputs.version }}
|
||||||
|
matrix: ${{ steps.resolve.outputs.matrix }}
|
||||||
|
steps:
|
||||||
|
- name: Resolve components and images
|
||||||
|
id: resolve
|
||||||
|
env:
|
||||||
|
COMPONENT: ${{ inputs.component }}
|
||||||
|
INPUT_VERSION: ${{ inputs.version }}
|
||||||
|
REPO_VARS: ${{ toJSON(vars) }}
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
version="${INPUT_VERSION#v}"
|
||||||
|
if [[ ! "$version" =~ ^[0-9]+\.[0-9]+\.[0-9]+$ ]]; then
|
||||||
|
echo "::error::Only stable x.y.z versions are certified, got '${INPUT_VERSION}'"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
# name, image repository, tag suffix (must match .goreleaser.yaml).
|
||||||
|
# Keep the names in sync with the workflow_dispatch options above.
|
||||||
|
components=(
|
||||||
|
"client-rootless ghcr.io/netbirdio/netbird -rootless-ubi"
|
||||||
|
"reverse-proxy ghcr.io/netbirdio/reverse-proxy -ubi"
|
||||||
|
"netbird-server ghcr.io/netbirdio/netbird-server -ubi"
|
||||||
|
)
|
||||||
|
matrix="[]"
|
||||||
|
missing=()
|
||||||
|
for c in "${components[@]}"; do
|
||||||
|
read -r name repo suffix <<< "$c"
|
||||||
|
[[ "$COMPONENT" == all || "$COMPONENT" == "$name" ]] || continue
|
||||||
|
var="REDHAT_CERT_ID_${name^^}"; var="${var//-/_}"
|
||||||
|
id="$(jq -r --arg v "$var" '.[$v] // empty' <<< "$REPO_VARS")"
|
||||||
|
if [[ -z "$id" ]]; then
|
||||||
|
missing+=("$var")
|
||||||
|
continue
|
||||||
|
fi
|
||||||
|
matrix="$(jq -c --arg n "$name" --arg t "${version}${suffix}" --arg r "${repo}:${version}${suffix}" --arg i "$id" \
|
||||||
|
'. + [{component: $n, tag: $t, ref: $r, component_id: $i}]' <<< "$matrix")"
|
||||||
|
done
|
||||||
|
if (( ${#missing[@]} )); then
|
||||||
|
echo "::error::Set these repository variables to the Partner Connect component IDs: ${missing[*]}"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
if [[ "$matrix" == "[]" ]]; then
|
||||||
|
echo "::error::No component to certify for '${COMPONENT}'"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
echo "Components to certify: ${matrix}"
|
||||||
|
echo "version=${version}" >> "$GITHUB_OUTPUT"
|
||||||
|
echo "matrix=${matrix}" >> "$GITHUB_OUTPUT"
|
||||||
|
|
||||||
|
certify:
|
||||||
|
name: "Certify ${{ matrix.component }} UBI image"
|
||||||
|
needs: resolve
|
||||||
|
runs-on: ubuntu-24.04
|
||||||
|
strategy:
|
||||||
|
fail-fast: false
|
||||||
|
matrix:
|
||||||
|
include: ${{ fromJSON(needs.resolve.outputs.matrix) }}
|
||||||
|
env:
|
||||||
|
PREFLIGHT_VERSION: "1.21.0"
|
||||||
|
# sha256 of preflight-linux-amd64 from the 1.21.0 GitHub release.
|
||||||
|
# Red Hat publishes no checksum file, so the value is pinned here.
|
||||||
|
PREFLIGHT_SHA256: "5e653135503c72f8702bbe31d7643197d12937c68086879133dd6b9650a9a449"
|
||||||
|
steps:
|
||||||
|
- name: Verify the multi-arch image is on ghcr.io
|
||||||
|
env:
|
||||||
|
IMAGE_REF: ${{ matrix.ref }}
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
docker buildx imagetools inspect "$IMAGE_REF" --raw > manifest.json
|
||||||
|
for arch in amd64 arm64; do
|
||||||
|
if ! jq -e --arg a "$arch" '.manifests[] | select(.platform.architecture == $a)' manifest.json > /dev/null; then
|
||||||
|
echo "::error::${IMAGE_REF} has no ${arch} manifest"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
done
|
||||||
|
echo "Manifest list for ${IMAGE_REF}:"
|
||||||
|
jq -r '.manifests[] | "\(.platform.os)/\(.platform.architecture) \(.digest)"' manifest.json
|
||||||
|
|
||||||
|
- name: Install preflight
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
curl -fsSL --proto '=https' --proto-redir '=https' -o preflight \
|
||||||
|
"https://github.com/redhat-openshift-ecosystem/openshift-preflight/releases/download/${PREFLIGHT_VERSION}/preflight-linux-amd64"
|
||||||
|
echo "${PREFLIGHT_SHA256} preflight" | sha256sum -c -
|
||||||
|
chmod +x preflight
|
||||||
|
./preflight --version
|
||||||
|
|
||||||
|
- name: Run preflight checks and submit to Red Hat
|
||||||
|
env:
|
||||||
|
IMAGE_REF: ${{ matrix.ref }}
|
||||||
|
PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
|
||||||
|
PFLT_CERTIFICATION_COMPONENT_ID: ${{ matrix.component_id }}
|
||||||
|
PFLT_ARTIFACTS: artifacts
|
||||||
|
PFLT_LOGFILE: artifacts/preflight.log
|
||||||
|
PFLT_LOGLEVEL: info
|
||||||
|
PFLT_JUNIT: "true"
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
# No --platform: preflight walks the manifest list and submits every
|
||||||
|
# architecture in one run, grouped under one manifest-list digest.
|
||||||
|
# preflight does not create the PFLT_LOGFILE directory, and --submit
|
||||||
|
# fails if the log file is missing.
|
||||||
|
mkdir -p artifacts
|
||||||
|
./preflight check container "$IMAGE_REF" --submit
|
||||||
|
|
||||||
|
- name: Fail if any check did not pass
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
shopt -s nullglob globstar
|
||||||
|
results=(artifacts/**/results.json)
|
||||||
|
if [[ ${#results[@]} -eq 0 ]]; then
|
||||||
|
echo "::error::preflight produced no results.json"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
status=0
|
||||||
|
for f in "${results[@]}"; do
|
||||||
|
arch="$(basename "$(dirname "$f")")"
|
||||||
|
passed="$(jq -r '.passed' "$f")"
|
||||||
|
failed="$(jq -r '[.results.failed[]?.name] | join(", ")' "$f")"
|
||||||
|
echo "${arch}: passed=${passed} ${failed:+failed checks: ${failed}}"
|
||||||
|
[[ "$passed" == "true" ]] || status=1
|
||||||
|
done
|
||||||
|
exit $status
|
||||||
|
|
||||||
|
- name: Upload preflight artifacts
|
||||||
|
if: always()
|
||||||
|
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
|
||||||
|
with:
|
||||||
|
name: redhat-preflight-${{ matrix.component }}-${{ needs.resolve.outputs.version }}
|
||||||
|
path: artifacts/
|
||||||
|
retention-days: 30
|
||||||
|
|
||||||
|
- name: Wait for Pyxis to mark both architectures certified
|
||||||
|
env:
|
||||||
|
TAG: ${{ matrix.tag }}
|
||||||
|
COMPONENT_ID: ${{ matrix.component_id }}
|
||||||
|
PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
# Filter on the tag server-side so older versions are found past the first page.
|
||||||
|
url="https://catalog.redhat.com/api/containers/v1/projects/certification/id/${COMPONENT_ID}/images?filter=repositories.tags.name==${TAG}&page_size=100"
|
||||||
|
for attempt in $(seq 1 20); do
|
||||||
|
certified="$(curl -fsS --proto '=https' --proto-redir '=https' -H "X-API-KEY: ${PFLT_PYXIS_API_TOKEN}" "$url" \
|
||||||
|
| jq -r --arg t "$TAG" '[.data[] | select(.repositories[]?.tags[]?.name == $t) | select(.certified == true) | .architecture] | unique | join(",")')"
|
||||||
|
echo "attempt ${attempt}: certified architectures for ${TAG}: ${certified:-none}"
|
||||||
|
if [[ "$certified" == "amd64,arm64" ]]; then
|
||||||
|
echo "Both architectures certified. Auto-publish is enabled on the component, so the catalog updates on its own."
|
||||||
|
exit 0
|
||||||
|
fi
|
||||||
|
sleep 30
|
||||||
|
done
|
||||||
|
echo "::error::Pyxis has not marked both architectures certified after 10 minutes. Check https://connect.redhat.com/component/view/${COMPONENT_ID}/images"
|
||||||
|
exit 1
|
||||||
@@ -69,7 +69,7 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
usesh: true
|
usesh: true
|
||||||
copyback: false
|
copyback: false
|
||||||
release: "15.0"
|
release: "15.1"
|
||||||
envs: "GO_VERSION"
|
envs: "GO_VERSION"
|
||||||
prepare: |
|
prepare: |
|
||||||
# Install required packages
|
# Install required packages
|
||||||
@@ -186,6 +186,22 @@ jobs:
|
|||||||
run: bash shared/management/http/api/generate.sh
|
run: bash shared/management/http/api/generate.sh
|
||||||
- name: check git status
|
- name: check git status
|
||||||
run: git --no-pager diff --exit-code
|
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
|
- name: Set up QEMU
|
||||||
uses: docker/setup-qemu-action@06116385d9baf250c9f4dcb4858b16962ea869c3 #v4.1.0
|
uses: docker/setup-qemu-action@06116385d9baf250c9f4dcb4858b16962ea869c3 #v4.1.0
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
@@ -225,14 +241,18 @@ jobs:
|
|||||||
uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2
|
uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2
|
||||||
with:
|
with:
|
||||||
version: ${{ env.GORELEASER_VER }}
|
version: ${{ env.GORELEASER_VER }}
|
||||||
args: release --clean ${{ env.flags }}
|
args: release --config .goreleaser.generated.yaml --clean ${{ env.flags }}
|
||||||
env:
|
env:
|
||||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||||
HOMEBREW_TAP_GITHUB_TOKEN: ${{ secrets.HOMEBREW_TAP_GITHUB_TOKEN }}
|
HOMEBREW_TAP_GITHUB_TOKEN: ${{ secrets.HOMEBREW_TAP_GITHUB_TOKEN }}
|
||||||
UPLOAD_DEBIAN_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }}
|
UPLOAD_DEBIAN_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }}
|
||||||
UPLOAD_YUM_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }}
|
UPLOAD_YUM_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }}
|
||||||
GPG_RPM_KEY_FILE: ${{ env.GPG_RPM_KEY_FILE }}
|
GPG_RPM_KEY_FILE: ${{ env.GPG_RPM_KEY_FILE }}
|
||||||
NFPM_NETBIRD_RPM_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
|
# One per nfpm id: GoReleaser looks the passphrase up as NFPM_<ID>_PASSPHRASE.
|
||||||
|
NFPM_NETBIRD_RPM_AMD64_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
|
||||||
|
NFPM_NETBIRD_RPM_ARM64_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
|
||||||
|
NFPM_NETBIRD_RPM_ARM_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
|
||||||
|
NFPM_NETBIRD_RPM_386_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
|
||||||
SKIP_PUBLISH: ${{ env.SKIP_PUBLISH }}
|
SKIP_PUBLISH: ${{ env.SKIP_PUBLISH }}
|
||||||
SKIP_DOCKER_PUSH: ${{ env.SKIP_DOCKER_PUSH }}
|
SKIP_DOCKER_PUSH: ${{ env.SKIP_DOCKER_PUSH }}
|
||||||
- name: Verify RPM signatures
|
- name: Verify RPM signatures
|
||||||
@@ -287,10 +307,17 @@ jobs:
|
|||||||
image_refs=()
|
image_refs=()
|
||||||
|
|
||||||
tag_and_push() {
|
tag_and_push() {
|
||||||
local src="$1" img_name tag dst
|
local src="$1" img_name tag dst variant=""
|
||||||
img_name="${src%%:*}"
|
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
|
for tag in $(resolve_tags); do
|
||||||
dst="${img_name}:${tag}"
|
dst="${img_name}:${tag}${variant}"
|
||||||
echo "Tagging ${src} -> ${dst}"
|
echo "Tagging ${src} -> ${dst}"
|
||||||
docker tag "$src" "$dst"
|
docker tag "$src" "$dst"
|
||||||
docker push "$dst"
|
docker push "$dst"
|
||||||
@@ -353,6 +380,24 @@ jobs:
|
|||||||
path: dist/netbird_darwin**
|
path: dist/netbird_darwin**
|
||||||
retention-days: 7
|
retention-days: 7
|
||||||
|
|
||||||
|
# Certify the UBI images in the Red Hat Ecosystem Catalog on stable tags.
|
||||||
|
# See redhat-certify.yml, which can also be run by hand for any released version.
|
||||||
|
redhat_certification:
|
||||||
|
name: "Red Hat"
|
||||||
|
needs: release
|
||||||
|
if: |
|
||||||
|
github.repository == 'netbirdio/netbird' &&
|
||||||
|
startsWith(github.ref, 'refs/tags/v') &&
|
||||||
|
!contains(github.ref_name, '-')
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
uses: ./.github/workflows/redhat-certify.yml
|
||||||
|
with:
|
||||||
|
component: all
|
||||||
|
version: ${{ github.ref_name }}
|
||||||
|
secrets:
|
||||||
|
PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
|
||||||
|
|
||||||
release_ui:
|
release_ui:
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
outputs:
|
outputs:
|
||||||
@@ -407,12 +452,12 @@ jobs:
|
|||||||
run: git --no-pager diff --exit-code
|
run: git --no-pager diff --exit-code
|
||||||
|
|
||||||
- name: Set up Node.js
|
- name: Set up Node.js
|
||||||
uses: actions/setup-node@v4
|
uses: actions/setup-node@v7
|
||||||
with:
|
with:
|
||||||
node-version: '22'
|
node-version: '22'
|
||||||
|
|
||||||
- name: Set up pnpm
|
- name: Set up pnpm
|
||||||
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
|
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
|
||||||
with:
|
with:
|
||||||
version: 11
|
version: 11
|
||||||
|
|
||||||
@@ -544,12 +589,12 @@ jobs:
|
|||||||
run: git --no-pager diff --exit-code
|
run: git --no-pager diff --exit-code
|
||||||
|
|
||||||
- name: Set up Node.js
|
- name: Set up Node.js
|
||||||
uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4
|
uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7.0.0
|
||||||
with:
|
with:
|
||||||
node-version: '22'
|
node-version: '22'
|
||||||
|
|
||||||
- name: Set up pnpm
|
- name: Set up pnpm
|
||||||
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
|
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
|
||||||
with:
|
with:
|
||||||
version: 11
|
version: 11
|
||||||
|
|
||||||
@@ -641,11 +686,11 @@ jobs:
|
|||||||
- name: check git status
|
- name: check git status
|
||||||
run: git --no-pager diff --exit-code
|
run: git --no-pager diff --exit-code
|
||||||
- name: Set up Node.js
|
- name: Set up Node.js
|
||||||
uses: actions/setup-node@v4
|
uses: actions/setup-node@v7
|
||||||
with:
|
with:
|
||||||
node-version: '22'
|
node-version: '22'
|
||||||
- name: Set up pnpm
|
- name: Set up pnpm
|
||||||
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
|
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
|
||||||
with:
|
with:
|
||||||
version: 11
|
version: 11
|
||||||
- name: Install wails3 CLI
|
- name: Install wails3 CLI
|
||||||
@@ -764,7 +809,7 @@ jobs:
|
|||||||
run: 7z x -o"${{ github.workspace }}/NSIS_Plugins" "${{ github.workspace }}/ShellExecAsUser_amd64-Unicode.7z"
|
run: 7z x -o"${{ github.workspace }}/NSIS_Plugins" "${{ github.workspace }}/ShellExecAsUser_amd64-Unicode.7z"
|
||||||
|
|
||||||
- name: Set up Go for wails3 CLI
|
- name: Set up Go for wails3 CLI
|
||||||
uses: actions/setup-go@v5
|
uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version-file: "go.mod"
|
go-version-file: "go.mod"
|
||||||
cache: false
|
cache: false
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -32,11 +32,12 @@ jobs:
|
|||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
- name: Set up Node.js
|
- name: Set up Node.js
|
||||||
uses: actions/setup-node@v4
|
uses: actions/setup-node@v7
|
||||||
with:
|
with:
|
||||||
node-version: "22"
|
node-version: "22"
|
||||||
|
|
||||||
# English (en) is the source of truth for translation keys; every other
|
# English (en) is the source of truth for translation keys. Locales declared
|
||||||
# locale declared in _index.json must carry the exact same key set.
|
# 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
|
- name: Check translation key parity
|
||||||
run: node client/ui/i18n/check-translations.mjs
|
run: node client/ui/i18n/check-translations.mjs
|
||||||
|
|||||||
@@ -35,3 +35,10 @@ vendor/
|
|||||||
/netbird
|
/netbird
|
||||||
client/netbird-electron/
|
client/netbird-electron/
|
||||||
management/server/types/testdata/
|
management/server/types/testdata/
|
||||||
|
|
||||||
|
# generated by chglog in the release workflow, embedded into the RPM
|
||||||
|
changelog.yml
|
||||||
|
|
||||||
|
# generated by rpm-provides.sh, the config GoReleaser actually runs
|
||||||
|
.goreleaser.generated.yaml
|
||||||
|
.chglog.yml
|
||||||
|
|||||||
+217
-6
@@ -40,6 +40,32 @@ builds:
|
|||||||
tags:
|
tags:
|
||||||
- load_wgnt_from_rsrc
|
- 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
|
- id: netbird-static
|
||||||
dir: client
|
dir: client
|
||||||
binary: netbird
|
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
|
- -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 }}"
|
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:
|
universal_binaries:
|
||||||
- id: netbird
|
- id: netbird
|
||||||
|
|
||||||
@@ -206,6 +254,10 @@ archives:
|
|||||||
builds:
|
builds:
|
||||||
- netbird-idp-migrate
|
- netbird-idp-migrate
|
||||||
name_template: "netbird-idp-migrate_{{ .Version }}_{{ .Os }}_{{ .Arch }}"
|
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:
|
nfpms:
|
||||||
- maintainer: Netbird <dev@netbird.io>
|
- maintainer: Netbird <dev@netbird.io>
|
||||||
@@ -223,23 +275,72 @@ nfpms:
|
|||||||
postinstall: "release_files/post_install.sh"
|
postinstall: "release_files/post_install.sh"
|
||||||
preremove: "release_files/pre_remove.sh"
|
preremove: "release_files/pre_remove.sh"
|
||||||
|
|
||||||
- maintainer: Netbird <dev@netbird.io>
|
- &netbird_rpm
|
||||||
|
maintainer: Netbird <dev@netbird.io>
|
||||||
description: Netbird client.
|
description: Netbird client.
|
||||||
homepage: https://netbird.io/
|
homepage: https://netbird.io/
|
||||||
license: BSD-3-Clause
|
license: BSD-3-Clause
|
||||||
vendor: NetBird
|
vendor: NetBird
|
||||||
id: netbird_rpm
|
id: netbird_rpm_amd64
|
||||||
bindir: /usr/bin
|
bindir: /usr/bin
|
||||||
builds:
|
ids:
|
||||||
- netbird
|
- netbird-rpm-amd64
|
||||||
formats:
|
formats:
|
||||||
- rpm
|
- rpm
|
||||||
|
# Red Hat certification (RPM Version Handling) requires rpmbuild's ISA
|
||||||
|
# provide, which nfpm does not emit. The version is filled in by the release job.
|
||||||
|
provides:
|
||||||
|
- "netbird(x86-64) = @RPM_EVR@"
|
||||||
|
# The client verifies TLS to management and signal against the system trust
|
||||||
|
# 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:
|
scripts:
|
||||||
postinstall: "release_files/post_install.sh"
|
postinstall: "release_files/post_install.sh"
|
||||||
preremove: "release_files/pre_remove.sh"
|
preremove: "release_files/pre_remove.sh"
|
||||||
rpm:
|
rpm:
|
||||||
|
summary: NetBird client
|
||||||
|
group: Applications/Internet
|
||||||
|
packager: NetBird <dev@netbird.io>
|
||||||
signature:
|
signature:
|
||||||
key_file: '{{ if index .Env "GPG_RPM_KEY_FILE" }}{{ .Env.GPG_RPM_KEY_FILE }}{{ end }}'
|
key_file: '{{ if index .Env "GPG_RPM_KEY_FILE" }}{{ .Env.GPG_RPM_KEY_FILE }}{{ end }}'
|
||||||
|
|
||||||
|
- <<: *netbird_rpm
|
||||||
|
id: netbird_rpm_arm64
|
||||||
|
ids:
|
||||||
|
- netbird-rpm-arm64
|
||||||
|
provides:
|
||||||
|
- "netbird(aarch-64) = @RPM_EVR@"
|
||||||
|
|
||||||
|
- <<: *netbird_rpm
|
||||||
|
id: netbird_rpm_arm
|
||||||
|
ids:
|
||||||
|
- netbird-rpm-arm
|
||||||
|
provides:
|
||||||
|
- "netbird(armv6hl-32) = @RPM_EVR@"
|
||||||
|
|
||||||
|
- <<: *netbird_rpm
|
||||||
|
id: netbird_rpm_386
|
||||||
|
ids:
|
||||||
|
- netbird-rpm-386
|
||||||
|
provides:
|
||||||
|
- "netbird(x86-32) = @RPM_EVR@"
|
||||||
dockers_v2:
|
dockers_v2:
|
||||||
- id: netbird
|
- id: netbird
|
||||||
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
|
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
|
||||||
@@ -289,6 +390,43 @@ dockers_v2:
|
|||||||
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
||||||
"org.opencontainers.image.source": "{{.GitURL}}"
|
"org.opencontainers.image.source": "{{.GitURL}}"
|
||||||
"maintainer": "dev@netbird.io"
|
"maintainer": "dev@netbird.io"
|
||||||
|
- id: 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
|
- id: relay
|
||||||
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
|
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
|
||||||
ids:
|
ids:
|
||||||
@@ -400,7 +538,7 @@ dockers_v2:
|
|||||||
tags:
|
tags:
|
||||||
- "{{ .Version }}"
|
- "{{ .Version }}"
|
||||||
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}latest{{ end }}"
|
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}latest{{ end }}"
|
||||||
dockerfile: upload-server/Dockerfile
|
dockerfile: upload-server/Dockerfile.release
|
||||||
platforms:
|
platforms:
|
||||||
- linux/amd64
|
- linux/amd64
|
||||||
- linux/arm64
|
- linux/arm64
|
||||||
@@ -434,6 +572,41 @@ dockers_v2:
|
|||||||
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
||||||
"org.opencontainers.image.source": "{{.GitURL}}"
|
"org.opencontainers.image.source": "{{.GitURL}}"
|
||||||
"maintainer": "dev@netbird.io"
|
"maintainer": "dev@netbird.io"
|
||||||
|
- id: 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
|
- id: netbird-proxy
|
||||||
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
|
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
|
||||||
ids:
|
ids:
|
||||||
@@ -456,6 +629,41 @@ dockers_v2:
|
|||||||
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
||||||
"org.opencontainers.image.source": "{{.GitURL}}"
|
"org.opencontainers.image.source": "{{.GitURL}}"
|
||||||
"maintainer": "dev@netbird.io"
|
"maintainer": "dev@netbird.io"
|
||||||
|
- id: proxy-ubi
|
||||||
|
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
|
||||||
|
ids:
|
||||||
|
- netbird-proxy
|
||||||
|
images:
|
||||||
|
- netbirdio/reverse-proxy
|
||||||
|
- ghcr.io/netbirdio/reverse-proxy
|
||||||
|
tags:
|
||||||
|
- "{{ .Version }}-ubi"
|
||||||
|
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}ubi-latest{{ end }}"
|
||||||
|
dockerfile: proxy/Dockerfile.ubi
|
||||||
|
platforms:
|
||||||
|
- linux/amd64
|
||||||
|
- linux/arm64
|
||||||
|
build_args:
|
||||||
|
VERSION: "{{ .Version }}"
|
||||||
|
RELEASE: "{{ .Timestamp }}"
|
||||||
|
hooks:
|
||||||
|
pre:
|
||||||
|
- cmd: 'sh 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:
|
brews:
|
||||||
- ids:
|
- ids:
|
||||||
@@ -488,7 +696,10 @@ uploads:
|
|||||||
- name: yum
|
- name: yum
|
||||||
skip: "{{ .Env.SKIP_PUBLISH }}"
|
skip: "{{ .Env.SKIP_PUBLISH }}"
|
||||||
ids:
|
ids:
|
||||||
- netbird_rpm
|
- netbird_rpm_amd64
|
||||||
|
- netbird_rpm_arm64
|
||||||
|
- netbird_rpm_arm
|
||||||
|
- netbird_rpm_386
|
||||||
mode: archive
|
mode: archive
|
||||||
target: https://pkgs.wiretrustee.com/yum/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }}
|
target: https://pkgs.wiretrustee.com/yum/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }}
|
||||||
username: dev@wiretrustee.com
|
username: dev@wiretrustee.com
|
||||||
|
|||||||
@@ -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 lint-all # full-repository lint, matches CI
|
||||||
make test-unit # host-safe unit tests, -tags devcert, no sudo
|
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 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
|
# Narrow runs
|
||||||
go test ./client/internal/dns/...
|
go test ./client/internal/dns/...
|
||||||
|
|||||||
@@ -1 +1,4 @@
|
|||||||
See [AGENTS.md](AGENTS.md) for the agent guidelines in this repository.
|
The agent guidelines live in [AGENTS.md](AGENTS.md). It is imported here so
|
||||||
|
every session loads it in full rather than following a pointer.
|
||||||
|
|
||||||
|
@AGENTS.md
|
||||||
|
|||||||
@@ -1,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.
|
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
|
BSD 3-Clause License
|
||||||
|
|||||||
@@ -23,8 +23,8 @@ lint-install: $(GOLANGCI_LINT)
|
|||||||
# Setup git hooks for all developers
|
# Setup git hooks for all developers
|
||||||
setup-hooks:
|
setup-hooks:
|
||||||
@git config core.hooksPath .githooks
|
@git config core.hooksPath .githooks
|
||||||
@chmod +x .githooks/pre-push
|
@chmod +x .githooks/pre-push .githooks/commit-msg
|
||||||
@echo "✅ Git hooks configured! Pre-push will now run 'make lint'"
|
@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).
|
# 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.
|
# Runs as a normal user with no sudo and leaves host networking untouched.
|
||||||
|
|||||||
@@ -115,6 +115,56 @@ export NETBIRD_DOMAIN=netbird.example.com; curl -fsSL https://github.com/netbird
|
|||||||
|
|
||||||
See a complete [architecture overview](https://docs.netbird.io/about-netbird/how-netbird-works#architecture) for details.
|
See a complete [architecture overview](https://docs.netbird.io/about-netbird/how-netbird-works#architecture) for details.
|
||||||
|
|
||||||
|
### Reporting bugs and requesting features
|
||||||
|
|
||||||
|
NetBird uses a discussion-first workflow. Bug reports and feature requests start in
|
||||||
|
[Discussions](https://github.com/netbirdio/netbird/discussions), not as issues.
|
||||||
|
|
||||||
|
| What you want to do | Where to go |
|
||||||
|
| --- | --- |
|
||||||
|
| Report a bug, regression, or unexpected behavior | [Issue Triage](https://github.com/netbirdio/netbird/discussions/new?category=issue-triage) |
|
||||||
|
| Request a feature or share an idea | [Ideas & Feature Requests](https://github.com/netbirdio/netbird/discussions/new?category=ideas-feature-requests) |
|
||||||
|
| Ask about setup, configuration, or self-hosting | [Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support) |
|
||||||
|
| Report a security vulnerability | [Security policy](https://github.com/netbirdio/netbird/security/policy), never a public thread |
|
||||||
|
|
||||||
|
Our team and maintainers triage discussions, ask follow-up questions, check for duplicates,
|
||||||
|
and reproduce bugs. Validated reports are promoted to issues. This keeps the issue tracker a clear
|
||||||
|
answer to one question: what is the team working on.
|
||||||
|
|
||||||
|
Please search existing discussions and issues first, including closed ones. If something similar
|
||||||
|
already exists, upvote it and add your details there instead of opening a duplicate.
|
||||||
|
|
||||||
|
For bug reports, include your NetBird version, operating system, deployment type (Cloud,
|
||||||
|
self-hosted, Kubernetes, or Docker), reproduction steps, expected and actual behavior, and a debug
|
||||||
|
bundle where relevant:
|
||||||
|
|
||||||
|
```shell
|
||||||
|
netbird version
|
||||||
|
netbird status -d -A
|
||||||
|
netbird debug for 1m -A -S -U
|
||||||
|
```
|
||||||
|
|
||||||
|
`-U` uploads the bundle and prints a file key you can paste instead of attaching the archive.
|
||||||
|
`-A` anonymizes the output, which matters on a public thread. It masks most identifying details
|
||||||
|
but is not full redaction, so read the bundle before posting it. Two levels are available:
|
||||||
|
|
||||||
|
| Level | How to select | What it masks |
|
||||||
|
| --- | --- | --- |
|
||||||
|
| `default` | `-A` / `--anonymize` | Public IP addresses, IPv6 ULA addresses, MAC addresses, and domains other than `netbird.io`, `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`. IPv4 private, CGNAT, and link-local ranges are kept |
|
||||||
|
| `strict` | `--anonymize-level strict` (implies `-A`) | The above, plus IPv4 private, CGNAT, and link-local ranges, peer names in front of `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`, and WireGuard public keys. Labels under `netbird.io` are kept, since it only hosts infrastructure |
|
||||||
|
|
||||||
|
See [collecting a debug bundle](https://docs.netbird.io/help/troubleshooting-client#debug-bundle)
|
||||||
|
and the [CLI reference](https://docs.netbird.io/get-started/cli#debug-for) for details.
|
||||||
|
|
||||||
|
See [How to use Discussions, Issues, and Pull Requests](https://github.com/netbirdio/netbird/discussions/6075)
|
||||||
|
for the full workflow, or [SUPPORT.md](SUPPORT.md) for a shorter version.
|
||||||
|
|
||||||
|
### Contributing
|
||||||
|
|
||||||
|
Contributions are welcome. Read [CONTRIBUTING.md](CONTRIBUTING.md) first. NetBird works ticket
|
||||||
|
first, anything that changes behavior needs an issue the team has agreed on before you open a pull
|
||||||
|
request.
|
||||||
|
|
||||||
### Community projects
|
### Community projects
|
||||||
- [NetBird installer script](https://github.com/physk/netbird-installer)
|
- [NetBird installer script](https://github.com/physk/netbird-installer)
|
||||||
- [netbird-tui](https://github.com/n0pashkov/netbird-tui) - terminal UI for managing NetBird peers, routes, and settings
|
- [netbird-tui](https://github.com/n0pashkov/netbird-tui) - terminal UI for managing NetBird peers, routes, and settings
|
||||||
|
|||||||
+121
@@ -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)
|
||||||
@@ -110,12 +110,12 @@ Two roles delegate Agent Network access without account-admin rights:
|
|||||||
read-only users, groups, peers, and account info (needed to build policies).
|
read-only users, groups, peers, and account info (needed to build policies).
|
||||||
Nothing else in the account.
|
Nothing else in the account.
|
||||||
- **`usage_viewer`** — the regular User baseline plus read on
|
- **`usage_viewer`** — the regular User baseline plus read on
|
||||||
`agent_network.usage` (the aggregated usage and cost overview) and read-only
|
`agent_network.usage` (the aggregated usage and cost overview) and
|
||||||
access to the resources the usage filters resolve against: users, groups,
|
`agent_network.logs` (the account-wide request-level access logs, which can
|
||||||
peers, and the provider list (connection config redacted — no upstream URLs
|
contain captured prompts), and read-only access to the resources those
|
||||||
or operator-supplied header values). No policies, and no account-wide
|
filters resolve against: users, groups, peers, and the provider list
|
||||||
request-level access logs; like any caller, it still reads its own requests
|
(connection config redacted — no upstream URLs or operator-supplied header
|
||||||
through the self-scoped endpoints below.
|
values). No policies, guardrails, budgets, or settings.
|
||||||
|
|
||||||
Every authenticated user, regardless of role, can read the caller-scoped
|
Every authenticated user, regardless of role, can read the caller-scoped
|
||||||
self-service endpoint `GET /api/agent-network/agent-config` (the endpoint, providers,
|
self-service endpoint `GET /api/agent-network/agent-config` (the endpoint, providers,
|
||||||
|
|||||||
+50
-31
@@ -3,56 +3,75 @@ package base62
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"math"
|
"math"
|
||||||
"strings"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
|
alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
|
||||||
base = uint32(len(alphabet))
|
base = uint32(len(alphabet))
|
||||||
|
maxBase62Digits = 6 // max number of digits required to encode MaxUint32
|
||||||
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
ErrEmptyString = fmt.Errorf("empty string")
|
||||||
|
ErrInvalidChar = fmt.Errorf("invalid character")
|
||||||
|
ErrOverflow = fmt.Errorf("integer overflow")
|
||||||
|
)
|
||||||
|
|
||||||
|
// Fixed-size arrays have better performance and lower memory overhead compared to maps for small static sets of data
|
||||||
|
var charToIndex [123]int8 // Assuming ASCII from '\0' - 'z'
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
for i := range charToIndex {
|
||||||
|
charToIndex[i] = -1
|
||||||
|
}
|
||||||
|
for i, c := range alphabet {
|
||||||
|
charToIndex[c] = int8(i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Encode encodes a uint32 value to a base62 string.
|
// Encode encodes a uint32 value to a base62 string.
|
||||||
func Encode(num uint32) string {
|
// The returned string will be between 1-6 characters long.
|
||||||
if num == 0 {
|
func Encode(n uint32) string {
|
||||||
return string(alphabet[0])
|
if n < base {
|
||||||
|
return string(alphabet[n])
|
||||||
|
}
|
||||||
|
// avoid dynamic memory usage for small, fixed size data
|
||||||
|
buf := [maxBase62Digits]byte{}
|
||||||
|
idx := len(buf)
|
||||||
|
|
||||||
|
for n > 0 {
|
||||||
|
idx--
|
||||||
|
buf[idx] = alphabet[n%base]
|
||||||
|
n /= base
|
||||||
}
|
}
|
||||||
|
|
||||||
var encoded strings.Builder
|
return string(buf[idx:])
|
||||||
|
|
||||||
for num > 0 {
|
|
||||||
remainder := num % base
|
|
||||||
encoded.WriteByte(alphabet[remainder])
|
|
||||||
num /= base
|
|
||||||
}
|
|
||||||
|
|
||||||
// Reverse the encoded string
|
|
||||||
encodedString := encoded.String()
|
|
||||||
reversed := reverse(encodedString)
|
|
||||||
return reversed
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Decode decodes a base62 string to a uint32 value.
|
// Decode decodes a base62 string to a uint32 value.
|
||||||
|
// Returns an error if the input string is empty, contains invalid characters,
|
||||||
|
// or would result in integer overflow.
|
||||||
func Decode(encoded string) (uint32, error) {
|
func Decode(encoded string) (uint32, error) {
|
||||||
|
if len(encoded) == 0 {
|
||||||
|
return 0, ErrEmptyString
|
||||||
|
}
|
||||||
var decoded uint32
|
var decoded uint32
|
||||||
strLen := len(encoded)
|
for _, char := range encoded {
|
||||||
|
index := int8(-1)
|
||||||
for i, char := range encoded {
|
if int(char) < len(charToIndex) {
|
||||||
index := strings.IndexRune(alphabet, char)
|
index = charToIndex[char]
|
||||||
|
}
|
||||||
if index < 0 {
|
if index < 0 {
|
||||||
return 0, fmt.Errorf("invalid character: %c", char)
|
return 0, fmt.Errorf("%w: %c", ErrInvalidChar, char)
|
||||||
|
}
|
||||||
|
// Add overflow check when calculating the decoded value to prevent silent overflow of uint32
|
||||||
|
if decoded > (math.MaxUint32-uint32(index))/base {
|
||||||
|
return 0, fmt.Errorf("%w: %s", ErrOverflow, encoded)
|
||||||
}
|
}
|
||||||
|
|
||||||
decoded += uint32(index) * uint32(math.Pow(float64(base), float64(strLen-i-1)))
|
decoded = decoded*base + uint32(index)
|
||||||
}
|
}
|
||||||
|
|
||||||
return decoded, nil
|
return decoded, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reverse a string.
|
|
||||||
func reverse(s string) string {
|
|
||||||
runes := []rune(s)
|
|
||||||
for i, j := 0, len(runes)-1; i < j; i, j = i+1, j-1 {
|
|
||||||
runes[i], runes[j] = runes[j], runes[i]
|
|
||||||
}
|
|
||||||
return string(runes)
|
|
||||||
}
|
|
||||||
|
|||||||
+50
-14
@@ -1,31 +1,67 @@
|
|||||||
package base62
|
package base62
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
|
"math"
|
||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestEncodeDecode(t *testing.T) {
|
func TestEncodeDecode(t *testing.T) {
|
||||||
tests := []struct {
|
testCases := []struct {
|
||||||
num uint32
|
input uint32
|
||||||
|
expected string
|
||||||
}{
|
}{
|
||||||
{0},
|
{0, "0"},
|
||||||
{1},
|
{1, "1"},
|
||||||
{42},
|
{5, "5"},
|
||||||
{12345},
|
{9, "9"},
|
||||||
{99999},
|
{10, "A"},
|
||||||
{123456789},
|
{42, "g"},
|
||||||
|
{61, "z"},
|
||||||
|
{62, "10"},
|
||||||
|
{'0', "m"},
|
||||||
|
{'9', "v"},
|
||||||
|
{'A', "13"},
|
||||||
|
{'Z', "1S"},
|
||||||
|
{'a', "1Z"},
|
||||||
|
{'z', "1y"},
|
||||||
|
{99999, "Q0t"},
|
||||||
|
{12345, "3D7"},
|
||||||
|
{123456789, "8M0kX"},
|
||||||
|
{math.MaxUint32, "4gfFC3"},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tc := range testCases {
|
||||||
encoded := Encode(tt.num)
|
encoded := Encode(tc.input)
|
||||||
|
if encoded != tc.expected {
|
||||||
|
t.Errorf("Encode(%d) = %s; want %s", tc.input, encoded, tc.expected)
|
||||||
|
}
|
||||||
decoded, err := Decode(encoded)
|
decoded, err := Decode(encoded)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Errorf("Decode error: %v", err)
|
t.Errorf("Expected error nil, got %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if decoded != tt.num {
|
if decoded != tc.input {
|
||||||
t.Errorf("Decode(%v) = %v, want %v", encoded, decoded, tt.num)
|
t.Errorf("Decode(%v) = %v, want %v", encoded, decoded, tc.input)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Decode handles empty string input with appropriate error
|
||||||
|
func TestDecodeEmptyString(t *testing.T) {
|
||||||
|
if _, err := Decode(""); !errors.Is(err, ErrEmptyString) {
|
||||||
|
t.Errorf("Expected error %v, got %v", ErrEmptyString, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDecodeOverflow(t *testing.T) {
|
||||||
|
if _, err := Decode("4gfFC4"); !errors.Is(err, ErrOverflow) {
|
||||||
|
t.Errorf("Expected error %v, got %v", ErrOverflow, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDecodeInvalid(t *testing.T) {
|
||||||
|
if _, err := Decode("/"); !errors.Is(err, ErrInvalidChar) {
|
||||||
|
t.Errorf("Expected error %v, got %v", ErrInvalidChar, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,45 @@
|
|||||||
|
FROM registry.access.redhat.com/ubi9/ubi-minimal@sha256:7fbeae18dc9476399f565e68255f602a3374ea8614ba3d14843565131a13ff93
|
||||||
|
|
||||||
|
ARG TARGETPLATFORM
|
||||||
|
ARG NETBIRD_BINARY=$TARGETPLATFORM/netbird
|
||||||
|
ARG VERSION=dev
|
||||||
|
ARG RELEASE=1
|
||||||
|
|
||||||
|
LABEL name="netbird-rootless" \
|
||||||
|
maintainer="NetBird <dev@netbird.io>" \
|
||||||
|
vendor="NetBird GmbH" \
|
||||||
|
version="${VERSION}" \
|
||||||
|
release="${RELEASE}" \
|
||||||
|
summary="NetBird Rootless Client" \
|
||||||
|
description="NetBird connects devices through an encrypted overlay using userspace networking without a TUN device or network administration capabilities."
|
||||||
|
|
||||||
|
RUN microdnf install -y bash ca-certificates && microdnf clean all
|
||||||
|
|
||||||
|
COPY --chmod=0555 client/netbird-entrypoint.sh /usr/local/bin/netbird-entrypoint.sh
|
||||||
|
COPY --chmod=0555 ${NETBIRD_BINARY} /usr/local/bin/netbird
|
||||||
|
COPY licenses/ /licenses/
|
||||||
|
# Only application storage is group-writable for arbitrary non-root UIDs.
|
||||||
|
# Runtime-created credentials keep the client's restrictive file modes.
|
||||||
|
RUN mkdir -p /var/lib/netbird && \
|
||||||
|
chown 1000:0 /var/lib/netbird && \
|
||||||
|
chmod 0770 /var/lib/netbird && \
|
||||||
|
chmod -R a+rX /licenses
|
||||||
|
|
||||||
|
WORKDIR /var/lib/netbird
|
||||||
|
USER 1000:0
|
||||||
|
|
||||||
|
ENV \
|
||||||
|
HOME="/var/lib/netbird" \
|
||||||
|
NETBIRD_BIN="/usr/local/bin/netbird" \
|
||||||
|
NB_USE_NETSTACK_MODE="true" \
|
||||||
|
NB_ENABLE_NETSTACK_LOCAL_FORWARDING="true" \
|
||||||
|
NB_CONFIG="/var/lib/netbird/config.json" \
|
||||||
|
NB_STATE_DIR="/var/lib/netbird" \
|
||||||
|
NB_DAEMON_ADDR="unix:///var/lib/netbird/netbird.sock" \
|
||||||
|
NB_LOG_FILE="console,/var/lib/netbird/client.log" \
|
||||||
|
NB_DISABLE_DNS="true" \
|
||||||
|
NB_ENABLE_CAPTURE="false" \
|
||||||
|
NB_ENTRYPOINT_SERVICE_TIMEOUT="30"
|
||||||
|
|
||||||
|
STOPSIGNAL SIGTERM
|
||||||
|
ENTRYPOINT ["/usr/local/bin/netbird-entrypoint.sh"]
|
||||||
+43
-11
@@ -9,6 +9,7 @@ import (
|
|||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"golang.org/x/exp/maps"
|
"golang.org/x/exp/maps"
|
||||||
@@ -90,13 +91,20 @@ type Client struct {
|
|||||||
connectClient *internal.ConnectClient
|
connectClient *internal.ConnectClient
|
||||||
config *profilemanager.Config
|
config *profilemanager.Config
|
||||||
cacheDir string
|
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.
|
// Identifies the running profile for the SSO login hint; see profile_state.go.
|
||||||
cfgPath string
|
cfgPath string
|
||||||
|
|
||||||
stateChangeMu sync.Mutex
|
stateChangeMu sync.Mutex
|
||||||
stateChangeSubID string
|
stateChangeSubID string
|
||||||
eventSub *peer.EventSubscription
|
// Closed to stop the watch goroutine from delivering buffered ticks to a
|
||||||
// Closed to stop the watch goroutines from delivering buffered items to a
|
|
||||||
// listener that has been removed or replaced. See stopStateChangeWatchLocked.
|
// listener that has been removed or replaced. See stopStateChangeWatchLocked.
|
||||||
stateChangeDone chan struct{}
|
stateChangeDone chan struct{}
|
||||||
|
|
||||||
@@ -178,6 +186,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
c.applyMDMOverlay(cfg)
|
||||||
c.recorder.UpdateManagementAddress(cfg.ManagementURL.String())
|
c.recorder.UpdateManagementAddress(cfg.ManagementURL.String())
|
||||||
c.recorder.UpdateRosenpass(cfg.RosenpassEnabled, cfg.RosenpassPermissive)
|
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,
|
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
|
||||||
internal.WithNetEvents(c.netMgr))
|
internal.WithNetEvents(c.netMgr))
|
||||||
c.setState(cfg, cacheDir, cfgFile, connectClient)
|
c.setState(cfg, cacheDir, cfgFile, connectClient)
|
||||||
|
connectClient.SetSyncResponsePersistence(true)
|
||||||
// This path runs the interactive SSO flow, so reaching here means the peer
|
// This path runs the interactive SSO flow, so reaching here means the peer
|
||||||
// is authenticated again — release the latch Status() reports from. Clear
|
// is authenticated again — release the latch Status() reports from. Clear
|
||||||
// only once the fresh connect client is installed: until then Status()
|
// only once the fresh connect client is installed: until then Status()
|
||||||
@@ -229,6 +239,7 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
c.applyMDMOverlay(cfg)
|
||||||
c.recorder.UpdateManagementAddress(cfg.ManagementURL.String())
|
c.recorder.UpdateManagementAddress(cfg.ManagementURL.String())
|
||||||
c.recorder.UpdateRosenpass(cfg.RosenpassEnabled, cfg.RosenpassPermissive)
|
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,
|
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
|
||||||
internal.WithNetEvents(c.netMgr))
|
internal.WithNetEvents(c.netMgr))
|
||||||
c.setState(cfg, cacheDir, cfgFile, connectClient)
|
c.setState(cfg, cacheDir, cfgFile, connectClient)
|
||||||
|
connectClient.SetSyncResponsePersistence(true)
|
||||||
return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir)
|
return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -316,6 +328,19 @@ func (c *Client) NotifyNetworkChange() {
|
|||||||
// or "strict"; strict also anonymizes internal IP ranges, peer names, and
|
// or "strict"; strict also anonymizes internal IP ranges, peer names, and
|
||||||
// WireGuard public keys, and implies anonymize.
|
// WireGuard public keys, and implies anonymize.
|
||||||
func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) {
|
func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) {
|
||||||
|
return c.debugBundle(platformFiles, anonymize, anonymizeLevel, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DebugBundleFile generates a debug bundle and returns the path of the zip in
|
||||||
|
// the cache directory instead of uploading it, so the app can hand the file to
|
||||||
|
// the user for inspection. The caller owns the file and removes it once done;
|
||||||
|
// the stale-bundle cleanup of later runs removes it only after a day.
|
||||||
|
// anonymize and anonymizeLevel behave as in DebugBundle.
|
||||||
|
func (c *Client) DebugBundleFile(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) {
|
||||||
|
return c.debugBundle(platformFiles, anonymize, anonymizeLevel, false)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) debugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string, upload bool) (string, error) {
|
||||||
cfg, cacheDir, cc := c.stateSnapshot()
|
cfg, cacheDir, cc := c.stateSnapshot()
|
||||||
|
|
||||||
// If the engine hasn't been started, load config from disk
|
// If the engine hasn't been started, load config from disk
|
||||||
@@ -327,9 +352,15 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("load config: %w", err)
|
return "", fmt.Errorf("load config: %w", err)
|
||||||
}
|
}
|
||||||
|
c.applyMDMOverlay(cfg)
|
||||||
cacheDir = platformFiles.CacheDir()
|
cacheDir = platformFiles.CacheDir()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Clear what an interrupted earlier run may have left in the cache before
|
||||||
|
// adding to it. Remote debug jobs write to the same directory, so anything
|
||||||
|
// younger than an hour is treated as possibly still in use.
|
||||||
|
debug.RemoveStaleBundles(cacheDir, time.Hour)
|
||||||
|
|
||||||
deps := debug.GeneratorDependencies{
|
deps := debug.GeneratorDependencies{
|
||||||
InternalConfig: cfg,
|
InternalConfig: cfg,
|
||||||
StatusRecorder: c.recorder,
|
StatusRecorder: c.recorder,
|
||||||
@@ -367,6 +398,9 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("generate debug bundle: %w", err)
|
return "", fmt.Errorf("generate debug bundle: %w", err)
|
||||||
}
|
}
|
||||||
|
if !upload {
|
||||||
|
return debug.ExportBundle(path)
|
||||||
|
}
|
||||||
defer func() {
|
defer func() {
|
||||||
if err := os.Remove(path); err != nil {
|
if err := os.Remove(path); err != nil {
|
||||||
log.Errorf("failed to remove debug bundle file: %v", err)
|
log.Errorf("failed to remove debug bundle file: %v", err)
|
||||||
@@ -463,6 +497,7 @@ func (c *Client) Networks() *NetworkArray {
|
|||||||
routesMap := routeManager.GetClientRoutesWithNetID()
|
routesMap := routeManager.GetClientRoutesWithNetID()
|
||||||
v6Merged := route.V6ExitMergeSet(routesMap)
|
v6Merged := route.V6ExitMergeSet(routesMap)
|
||||||
resolvedDomains := c.recorder.GetResolvedDomainsStates()
|
resolvedDomains := c.recorder.GetResolvedDomainsStates()
|
||||||
|
activeRoutePeers := c.recorder.GetActiveRoutePeers()
|
||||||
|
|
||||||
networkArray := &NetworkArray{
|
networkArray := &NetworkArray{
|
||||||
items: make([]Network, 0),
|
items: make([]Network, 0),
|
||||||
@@ -476,7 +511,7 @@ func (c *Client) Networks() *NetworkArray {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged)
|
network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged, activeRoutePeers)
|
||||||
if network == nil {
|
if network == nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -485,14 +520,14 @@ func (c *Client) Networks() *NetworkArray {
|
|||||||
return networkArray
|
return networkArray
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}) *Network {
|
func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}, activeRoutePeers map[route.HAUniqueID]string) *Network {
|
||||||
r := routes[0]
|
r := routes[0]
|
||||||
netStr := r.Network.String()
|
netStr := r.Network.String()
|
||||||
if r.IsDynamic() {
|
if r.IsDynamic() {
|
||||||
netStr = r.Domains.SafeString()
|
netStr = r.Domains.SafeString()
|
||||||
}
|
}
|
||||||
|
|
||||||
routePeer, err := c.findBestRoutePeer(routes)
|
routePeer, err := c.findBestRoutePeer(routes, activeRoutePeers)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Errorf("could not get peer info for route %s: %v", id, err)
|
log.Errorf("could not get peer info for route %s: %v", id, err)
|
||||||
return nil
|
return nil
|
||||||
@@ -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
|
// findBestRoutePeer returns the peer actively routing traffic for the given
|
||||||
// HA route group. Falls back to the first connected peer, then the first peer.
|
// HA route group. Falls back to the first connected peer, then the first peer.
|
||||||
func (c *Client) findBestRoutePeer(routes []*route.Route) (peer.State, error) {
|
func (c *Client) findBestRoutePeer(routes []*route.Route, activeRoutePeers map[route.HAUniqueID]string) (peer.State, error) {
|
||||||
netStr := routes[0].Network.String()
|
if peerKey, ok := activeRoutePeers[routes[0].GetHAUniqueID()]; ok {
|
||||||
|
if p, err := c.recorder.GetPeer(peerKey); err == nil {
|
||||||
fullStatus := c.recorder.GetFullStatus()
|
|
||||||
for _, p := range fullStatus.Peers {
|
|
||||||
if _, ok := p.GetRoutes()[netStr]; ok {
|
|
||||||
return p, nil
|
return p, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,52 @@
|
|||||||
|
//go:build android
|
||||||
|
|
||||||
|
package android
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
"github.com/netbirdio/netbird/client/mdm"
|
||||||
|
)
|
||||||
|
|
||||||
|
type mdmSource struct {
|
||||||
|
loader *mdm.Loader
|
||||||
|
detector *mdm.ChangeDetector
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetMDMPolicyFetcher registers the native-provided MDM policy fetcher on
|
||||||
|
// this Client; passing nil disables MDM enforcement.
|
||||||
|
func (c *Client) SetMDMPolicyFetcher(p PolicyFetcher) {
|
||||||
|
loader := loaderFor(p)
|
||||||
|
c.mdmSource.Store(&mdmSource{loader: loader, detector: mdm.NewChangeDetector(loader)})
|
||||||
|
}
|
||||||
|
|
||||||
|
// HasMDMPolicyChanged re-reads the managed configuration and reports whether
|
||||||
|
// it changed since the last observation; call it from the native OS-change
|
||||||
|
// notification and restart the engine only on true.
|
||||||
|
func (c *Client) HasMDMPolicyChanged() bool {
|
||||||
|
src := c.mdmSource.Load()
|
||||||
|
if src == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return src.detector.Changed()
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetRestrictionsJSON returns the UI enforcement snapshot derived from the
|
||||||
|
// active MDM policy, in the JSON shape shared with the desktop frontend.
|
||||||
|
func (c *Client) GetRestrictionsJSON() (string, error) {
|
||||||
|
return mdm.BuildRestrictions(c.mdmLoader().Load()).JSON()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) applyMDMOverlay(cfg *profilemanager.Config) {
|
||||||
|
loader := c.mdmLoader()
|
||||||
|
if cfg == nil || loader == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cfg.ApplyMDMPolicy(loader.Load())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) mdmLoader() *mdm.Loader {
|
||||||
|
if src := c.mdmSource.Load(); src != nil {
|
||||||
|
return src.loader
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+56
-18
@@ -8,6 +8,7 @@ import (
|
|||||||
|
|
||||||
"github.com/netbirdio/netbird/client/internal/auth"
|
"github.com/netbirdio/netbird/client/internal/auth"
|
||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
"github.com/netbirdio/netbird/client/mdm"
|
||||||
"github.com/netbirdio/netbird/client/mobile"
|
"github.com/netbirdio/netbird/client/mobile"
|
||||||
"github.com/netbirdio/netbird/client/system"
|
"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
|
// 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
|
// 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".
|
// management stream rejects it with "no peer auth method provided".
|
||||||
func NewAuth(cfgPath string, mgmURL string) (*Auth, error) {
|
//
|
||||||
inputCfg := profilemanager.ConfigInput{
|
// Auth is constructed under the active MDM policy: the policy is overlaid on
|
||||||
ConfigPath: cfgPath,
|
// the resolved config so the login runs against the enforced values, while
|
||||||
ManagementURL: mgmURL,
|
// 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)
|
cfg, err := profilemanager.UpdateOrCreateConfig(inputCfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
cfg.ApplyMDMPolicy(policy)
|
||||||
|
|
||||||
return &Auth{
|
return &Auth{
|
||||||
ctx: context.Background(),
|
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.
|
// SaveConfigIfSSOSupported reports whether the management server supports SSO; the config is already persisted by NewAuth.
|
||||||
// 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.
|
|
||||||
func (a *Auth) SaveConfigIfSSOSupported(listener SSOListener) {
|
func (a *Auth) SaveConfigIfSSOSupported(listener SSOListener) {
|
||||||
go func() {
|
go func() {
|
||||||
sso, err := a.saveConfigIfSSOSupported()
|
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)
|
return false, fmt.Errorf("failed to check SSO support: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !supportsSSO {
|
return supportsSSO, nil
|
||||||
return false, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
err = profilemanager.WriteOutConfig(a.cfgPath, a.config)
|
// LoginWithSetupKeyAndSaveConfig registers the peer with the setup key; the config is already persisted by NewAuth.
|
||||||
return true, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// LoginWithSetupKeyAndSaveConfig test the connectivity with the management server with the setup key.
|
|
||||||
func (a *Auth) LoginWithSetupKeyAndSaveConfig(resultListener ErrListener, setupKey string, deviceName string) {
|
func (a *Auth) LoginWithSetupKeyAndSaveConfig(resultListener ErrListener, setupKey string, deviceName string) {
|
||||||
go func() {
|
go func() {
|
||||||
err := a.loginWithSetupKeyAndSaveConfig(setupKey, deviceName)
|
err := a.loginWithSetupKeyAndSaveConfig(setupKey, deviceName)
|
||||||
@@ -134,8 +136,7 @@ func (a *Auth) loginWithSetupKeyAndSaveConfig(setupKey string, deviceName string
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("login failed: %v", err)
|
return fmt.Errorf("login failed: %v", err)
|
||||||
}
|
}
|
||||||
|
return nil
|
||||||
return profilemanager.WriteOutConfig(a.cfgPath, a.config)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Login try register the client on the server
|
// 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) {
|
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 {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to get OAuth flow: %v", err)
|
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.
|
// profileLoginHint returns the stored account email for the profile at cfgPath.
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ import (
|
|||||||
func TestNewAuth_ReusesPersistedIdentity(t *testing.T) {
|
func TestNewAuth_ReusesPersistedIdentity(t *testing.T) {
|
||||||
cfgPath := filepath.Join(t.TempDir(), "config.json")
|
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 {
|
if err != nil {
|
||||||
t.Fatalf("first NewAuth: %v", err)
|
t.Fatalf("first NewAuth: %v", err)
|
||||||
}
|
}
|
||||||
@@ -24,7 +24,7 @@ func TestNewAuth_ReusesPersistedIdentity(t *testing.T) {
|
|||||||
t.Fatal("first NewAuth produced no private key")
|
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 {
|
if err != nil {
|
||||||
t.Fatalf("second NewAuth: %v", err)
|
t.Fatalf("second NewAuth: %v", err)
|
||||||
}
|
}
|
||||||
@@ -38,7 +38,7 @@ func TestNewAuth_ReusesPersistedIdentity(t *testing.T) {
|
|||||||
func TestNewAuth_CreatesConfigWhenAbsent(t *testing.T) {
|
func TestNewAuth_CreatesConfigWhenAbsent(t *testing.T) {
|
||||||
cfgPath := filepath.Join(t.TempDir(), "config.json")
|
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 {
|
if err != nil {
|
||||||
t.Fatalf("NewAuth: %v", err)
|
t.Fatalf("NewAuth: %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -1,12 +1,16 @@
|
|||||||
package android
|
package android
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
"github.com/netbirdio/netbird/client/mdm"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Preferences exports a subset of the internal config for gomobile
|
// Preferences exports a subset of the internal config for gomobile
|
||||||
type Preferences struct {
|
type Preferences struct {
|
||||||
configInput profilemanager.ConfigInput
|
configInput profilemanager.ConfigInput
|
||||||
|
mdmLoader atomic.Pointer[mdm.Loader]
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewPreferences creates a new Preferences instance
|
// NewPreferences creates a new Preferences instance
|
||||||
@@ -14,20 +18,39 @@ func NewPreferences(configPath string) *Preferences {
|
|||||||
ci := profilemanager.ConfigInput{
|
ci := profilemanager.ConfigInput{
|
||||||
ConfigPath: configPath,
|
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
|
// GetManagementURL reads URL from config file
|
||||||
func (p *Preferences) GetManagementURL() (string, error) {
|
func (p *Preferences) GetManagementURL() (string, error) {
|
||||||
|
if v, ok := p.policy().GetString(mdm.KeyManagementURL); ok {
|
||||||
|
return mdm.CanonicalURL(v), nil
|
||||||
|
}
|
||||||
if p.configInput.ManagementURL != "" {
|
if p.configInput.ManagementURL != "" {
|
||||||
return p.configInput.ManagementURL, nil
|
return p.configInput.ManagementURL, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
|
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
return cfg.ManagementURL.String(), err
|
return cfg.ManagementURL.String(), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetManagementURL stores the given URL and waits for commit
|
// SetManagementURL stores the given URL and waits for commit
|
||||||
@@ -41,7 +64,7 @@ func (p *Preferences) GetAdminURL() (string, error) {
|
|||||||
return p.configInput.AdminURL, nil
|
return p.configInput.AdminURL, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
|
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
@@ -53,17 +76,21 @@ func (p *Preferences) SetAdminURL(url string) {
|
|||||||
p.configInput.AdminURL = url
|
p.configInput.AdminURL = url
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetPreSharedKey reads pre-shared key from config file
|
// HasPreSharedKey reports whether a pre-shared key is staged, persisted, or
|
||||||
func (p *Preferences) GetPreSharedKey() (string, error) {
|
// 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 {
|
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 {
|
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
|
// 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
|
// GetRosenpassEnabled reads Rosenpass enabled status from config file
|
||||||
func (p *Preferences) GetRosenpassEnabled() (bool, error) {
|
func (p *Preferences) GetRosenpassEnabled() (bool, error) {
|
||||||
|
if v, ok := p.policy().GetBool(mdm.KeyRosenpassEnabled); ok {
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
if p.configInput.RosenpassEnabled != nil {
|
if p.configInput.RosenpassEnabled != nil {
|
||||||
return *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 {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -96,11 +126,14 @@ func (p *Preferences) SetRosenpassPermissive(permissive bool) {
|
|||||||
|
|
||||||
// GetRosenpassPermissive reads Rosenpass permissive setting from config file
|
// GetRosenpassPermissive reads Rosenpass permissive setting from config file
|
||||||
func (p *Preferences) GetRosenpassPermissive() (bool, error) {
|
func (p *Preferences) GetRosenpassPermissive() (bool, error) {
|
||||||
|
if v, ok := p.policy().GetBool(mdm.KeyRosenpassPermissive); ok {
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
if p.configInput.RosenpassPermissive != nil {
|
if p.configInput.RosenpassPermissive != nil {
|
||||||
return *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 {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -109,11 +142,14 @@ func (p *Preferences) GetRosenpassPermissive() (bool, error) {
|
|||||||
|
|
||||||
// GetDisableClientRoutes reads disable client routes setting from config file
|
// GetDisableClientRoutes reads disable client routes setting from config file
|
||||||
func (p *Preferences) GetDisableClientRoutes() (bool, error) {
|
func (p *Preferences) GetDisableClientRoutes() (bool, error) {
|
||||||
|
if v, ok := p.policy().GetBool(mdm.KeyDisableClientRoutes); ok {
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
if p.configInput.DisableClientRoutes != nil {
|
if p.configInput.DisableClientRoutes != nil {
|
||||||
return *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 {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -127,11 +163,14 @@ func (p *Preferences) SetDisableClientRoutes(disable bool) {
|
|||||||
|
|
||||||
// GetDisableServerRoutes reads disable server routes setting from config file
|
// GetDisableServerRoutes reads disable server routes setting from config file
|
||||||
func (p *Preferences) GetDisableServerRoutes() (bool, error) {
|
func (p *Preferences) GetDisableServerRoutes() (bool, error) {
|
||||||
|
if v, ok := p.policy().GetBool(mdm.KeyDisableServerRoutes); ok {
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
if p.configInput.DisableServerRoutes != nil {
|
if p.configInput.DisableServerRoutes != nil {
|
||||||
return *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 {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -149,7 +188,7 @@ func (p *Preferences) GetDisableDNS() (bool, error) {
|
|||||||
return *p.configInput.DisableDNS, nil
|
return *p.configInput.DisableDNS, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
|
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -167,7 +206,7 @@ func (p *Preferences) GetDisableFirewall() (bool, error) {
|
|||||||
return *p.configInput.DisableFirewall, nil
|
return *p.configInput.DisableFirewall, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
|
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -181,11 +220,14 @@ func (p *Preferences) SetDisableFirewall(disable bool) {
|
|||||||
|
|
||||||
// GetServerSSHAllowed reads server SSH allowed setting from config file
|
// GetServerSSHAllowed reads server SSH allowed setting from config file
|
||||||
func (p *Preferences) GetServerSSHAllowed() (bool, error) {
|
func (p *Preferences) GetServerSSHAllowed() (bool, error) {
|
||||||
|
if v, ok := p.policy().GetBool(mdm.KeyAllowServerSSH); ok {
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
if p.configInput.ServerSSHAllowed != nil {
|
if p.configInput.ServerSSHAllowed != nil {
|
||||||
return *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 {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -207,7 +249,7 @@ func (p *Preferences) GetEnableSSHRoot() (bool, error) {
|
|||||||
return *p.configInput.EnableSSHRoot, nil
|
return *p.configInput.EnableSSHRoot, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
|
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -229,7 +271,7 @@ func (p *Preferences) GetEnableSSHSFTP() (bool, error) {
|
|||||||
return *p.configInput.EnableSSHSFTP, nil
|
return *p.configInput.EnableSSHSFTP, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
|
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -251,7 +293,7 @@ func (p *Preferences) GetEnableSSHLocalPortForwarding() (bool, error) {
|
|||||||
return *p.configInput.EnableSSHLocalPortForwarding, nil
|
return *p.configInput.EnableSSHLocalPortForwarding, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
|
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -273,7 +315,7 @@ func (p *Preferences) GetEnableSSHRemotePortForwarding() (bool, error) {
|
|||||||
return *p.configInput.EnableSSHRemotePortForwarding, nil
|
return *p.configInput.EnableSSHRemotePortForwarding, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
|
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -291,11 +333,14 @@ func (p *Preferences) SetEnableSSHRemotePortForwarding(enabled bool) {
|
|||||||
|
|
||||||
// GetBlockInbound reads block inbound setting from config file
|
// GetBlockInbound reads block inbound setting from config file
|
||||||
func (p *Preferences) GetBlockInbound() (bool, error) {
|
func (p *Preferences) GetBlockInbound() (bool, error) {
|
||||||
|
if v, ok := p.policy().GetBool(mdm.KeyBlockInbound); ok {
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
if p.configInput.BlockInbound != nil {
|
if p.configInput.BlockInbound != nil {
|
||||||
return *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 {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -313,7 +358,7 @@ func (p *Preferences) GetDisableIPv6() (bool, error) {
|
|||||||
return *p.configInput.DisableIPv6, nil
|
return *p.configInput.DisableIPv6, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
|
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -327,18 +372,20 @@ func (p *Preferences) SetDisableIPv6(disable bool) {
|
|||||||
|
|
||||||
// GetRemoteJobsAllowed reads the remote jobs opt-in from config file
|
// GetRemoteJobsAllowed reads the remote jobs opt-in from config file
|
||||||
func (p *Preferences) GetRemoteJobsAllowed() (bool, error) {
|
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
|
return *p.configInput.RemoteJobsAllowed, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
|
cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
|
cfg.ApplyMDMPolicy(policy)
|
||||||
if cfg.RemoteJobsAllowed == nil {
|
if cfg.RemoteJobsAllowed == nil {
|
||||||
return false, nil
|
return false, nil
|
||||||
}
|
}
|
||||||
return *cfg.RemoteJobsAllowed, err
|
return *cfg.RemoteJobsAllowed, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetRemoteJobsAllowed stores the given value and waits for commit
|
// 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
|
// Commit writes out the changes to the config file
|
||||||
func (p *Preferences) Commit() error {
|
func (p *Preferences) Commit() error {
|
||||||
|
if err := profilemanager.CheckMDMConflicts(p.configInput, p.policy()); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
_, err := profilemanager.UpdateOrCreateConfig(p.configInput)
|
_, err := profilemanager.UpdateOrCreateConfig(p.configInput)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -28,14 +28,13 @@ func TestPreferences_DefaultValues(t *testing.T) {
|
|||||||
t.Errorf("invalid default management url: %s", defaultVar)
|
t.Errorf("invalid default management url: %s", defaultVar)
|
||||||
}
|
}
|
||||||
|
|
||||||
var preSharedKey string
|
hasPSK, err := p.HasPreSharedKey()
|
||||||
preSharedKey, err = p.GetPreSharedKey()
|
|
||||||
if err != nil {
|
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 != "" {
|
if hasPSK {
|
||||||
t.Errorf("invalid preshared key: %s", preSharedKey)
|
t.Errorf("unexpected preshared key presence on fresh config")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -65,13 +64,13 @@ func TestPreferences_ReadUncommitedValues(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
p.SetPreSharedKey(exampleString)
|
p.SetPreSharedKey(exampleString)
|
||||||
resp, err = p.GetPreSharedKey()
|
hasPSK, err := p.HasPreSharedKey()
|
||||||
if err != nil {
|
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 {
|
if !hasPSK {
|
||||||
t.Errorf("unexpected preshared key: %s", resp)
|
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)
|
t.Errorf("unexpected management url: %s", resp)
|
||||||
}
|
}
|
||||||
|
|
||||||
resp, err = p.GetPreSharedKey()
|
hasPSK, err := p.HasPreSharedKey()
|
||||||
if err != nil {
|
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 {
|
if !hasPSK {
|
||||||
t.Errorf("unexpected preshared key: %s", resp)
|
t.Errorf("expected preshared key presence after commit")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -54,6 +54,12 @@ func NewProfileManager(configDir string) *ProfileManager {
|
|||||||
return &ProfileManager{impl: mobile.NewProfileManager(configDir, androidUsername)}
|
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,
|
// ListProfiles returns all available profiles, including the default profile,
|
||||||
// with their active status set.
|
// with their active status set.
|
||||||
func (pm *ProfileManager) ListProfiles() (*ProfileArray, error) {
|
func (pm *ProfileManager) ListProfiles() (*ProfileArray, error) {
|
||||||
|
|||||||
+16
-85
@@ -6,13 +6,8 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/internal"
|
"github.com/netbirdio/netbird/client/internal"
|
||||||
"github.com/netbirdio/netbird/client/internal/auth"
|
"github.com/netbirdio/netbird/client/internal/auth"
|
||||||
"github.com/netbirdio/netbird/client/internal/auth/sessionwatch"
|
|
||||||
"github.com/netbirdio/netbird/client/internal/peer"
|
|
||||||
cProto "github.com/netbirdio/netbird/client/proto"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// StateChangeListener receives client state notifications.
|
// StateChangeListener receives client state notifications.
|
||||||
@@ -21,16 +16,11 @@ import (
|
|||||||
// changed: connection state, the run-loop status label (e.g. NeedsLogin) or
|
// changed: connection state, the run-loop status label (e.g. NeedsLogin) or
|
||||||
// the session deadline. It mirrors the daemon's SubscribeStatus stream
|
// the session deadline. It mirrors the daemon's SubscribeStatus stream
|
||||||
// trigger — on each signal the consumer pulls the fresh values via
|
// trigger — on each signal the consumer pulls the fresh values via
|
||||||
// Status() / SessionExpiresAtUnix().
|
// Status() / SessionExpiresAtUnix(). The engine arms no expiry-warning
|
||||||
//
|
// timers on Android; the app schedules the warnings from the deadline it
|
||||||
// OnSessionExpiring forwards the engine's session-expiry warnings, fired at
|
// reads here.
|
||||||
// sessionwatch.WarningLead before the deadline and again at FinalWarningLead
|
|
||||||
// (finalWarning true). The second one is suppressed when the user dismissed
|
|
||||||
// the first via DismissSessionWarning. The daemon turns the same events into
|
|
||||||
// its tray notification.
|
|
||||||
type StateChangeListener interface {
|
type StateChangeListener interface {
|
||||||
OnStateChanged()
|
OnStateChanged()
|
||||||
OnSessionExpiring(expiresAtUnix int64, leadMinutes int64, finalWarning bool)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Status returns the connect run-loop's status label — the same value the
|
// Status returns the connect run-loop's status label — the same value the
|
||||||
@@ -110,11 +100,11 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Both subscriptions are buffered (one pending tick, ten pending events),
|
// The subscription is buffered (one pending tick), so unsubscribing is
|
||||||
// so unsubscribing is not enough to stop callbacks: the loops would drain
|
// not enough to stop callbacks: the loop would drain what is already
|
||||||
// what is already queued and deliver it to a listener the caller has
|
// queued and deliver it to a listener the caller has already removed or
|
||||||
// already removed or replaced. Gate every callback on this registration's
|
// replaced. Gate every callback on this registration's own signal, which
|
||||||
// own signal, which is closed before unsubscribing.
|
// is closed before unsubscribing.
|
||||||
done := make(chan struct{})
|
done := make(chan struct{})
|
||||||
c.stateChangeDone = done
|
c.stateChangeDone = done
|
||||||
|
|
||||||
@@ -133,9 +123,6 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) {
|
|||||||
listener.OnStateChanged()
|
listener.OnStateChanged()
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
c.eventSub = c.recorder.SubscribeToEvents()
|
|
||||||
go watchSessionWarnings(c.eventSub, listener, done)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// RemoveStateChangeListener unregisters the state notification listener.
|
// RemoveStateChangeListener unregisters the state notification listener.
|
||||||
@@ -145,21 +132,6 @@ func (c *Client) RemoveStateChangeListener() {
|
|||||||
c.stopStateChangeWatchLocked()
|
c.stopStateChangeWatchLocked()
|
||||||
}
|
}
|
||||||
|
|
||||||
// DismissSessionWarning records the user's "Dismiss" on the first expiry
|
|
||||||
// warning and suppresses the final one for the current deadline. A refreshed
|
|
||||||
// deadline re-arms both. No-op while the engine is not running.
|
|
||||||
func (c *Client) DismissSessionWarning() {
|
|
||||||
cc := c.getConnectClient()
|
|
||||||
if cc == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
engine := cc.Engine()
|
|
||||||
if engine == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
engine.DismissSessionWarning()
|
|
||||||
}
|
|
||||||
|
|
||||||
// ExtendAuthSession runs the interactive SSO flow to obtain a fresh JWT and
|
// ExtendAuthSession runs the interactive SSO flow to obtain a fresh JWT and
|
||||||
// asks the management server to extend the session deadline. The tunnel is
|
// asks the management server to extend the session deadline. The tunnel is
|
||||||
// untouched: no resync, no reconnect. Async; the result arrives on the
|
// untouched: no resync, no reconnect. Async; the result arrives on the
|
||||||
@@ -201,8 +173,8 @@ func (c *Client) CancelExtendAuthSession() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) stopStateChangeWatchLocked() {
|
func (c *Client) stopStateChangeWatchLocked() {
|
||||||
// Signal first, unsubscribe second: closing the channels only stops new
|
// Signal first, unsubscribe second: closing the channel only stops new
|
||||||
// items, and the loops would still hand whatever is buffered to a listener
|
// items, and the loop would still hand whatever is buffered to a listener
|
||||||
// that is no longer registered.
|
// that is no longer registered.
|
||||||
if c.stateChangeDone != nil {
|
if c.stateChangeDone != nil {
|
||||||
close(c.stateChangeDone)
|
close(c.stateChangeDone)
|
||||||
@@ -212,49 +184,6 @@ func (c *Client) stopStateChangeWatchLocked() {
|
|||||||
c.recorder.UnsubscribeFromStateChanges(c.stateChangeSubID)
|
c.recorder.UnsubscribeFromStateChanges(c.stateChangeSubID)
|
||||||
c.stateChangeSubID = ""
|
c.stateChangeSubID = ""
|
||||||
}
|
}
|
||||||
if c.eventSub != nil {
|
|
||||||
// Closes the channel, which ends watchSessionWarnings.
|
|
||||||
c.recorder.UnsubscribeFromEvents(c.eventSub)
|
|
||||||
c.eventSub = nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// watchSessionWarnings forwards the engine's session-expiry warnings to the
|
|
||||||
// listener. The event stream also carries unrelated traffic — network-map
|
|
||||||
// updates on every sync, DNS and route errors — so everything but an
|
|
||||||
// AUTHENTICATION event carrying the session-warning marker is dropped. Exits
|
|
||||||
// when the subscription is closed by UnsubscribeFromEvents, or earlier when
|
|
||||||
// done is closed — the stream buffers up to ten events, and a deregistered
|
|
||||||
// listener must not receive the ones already queued.
|
|
||||||
func watchSessionWarnings(sub *peer.EventSubscription, listener StateChangeListener, done <-chan struct{}) {
|
|
||||||
for ev := range sub.Events() {
|
|
||||||
select {
|
|
||||||
case <-done:
|
|
||||||
return
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
if ev.GetCategory() != cProto.SystemEvent_AUTHENTICATION {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
meta := ev.GetMetadata()
|
|
||||||
if meta[sessionwatch.MetaSessionWarning] != "true" {
|
|
||||||
// Other AUTHENTICATION events exist (e.g. a deadline rejected as
|
|
||||||
// out of range); they carry no warning marker.
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
deadline, err := sessionwatch.ParseExpiresAt(meta[sessionwatch.MetaSessionExpiresAt])
|
|
||||||
if err != nil {
|
|
||||||
log.Warnf("session warning event with unparsable deadline: %v", err)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
lead, err := sessionwatch.ParseLeadMinutes(meta[sessionwatch.MetaSessionLeadMinutes])
|
|
||||||
if err != nil {
|
|
||||||
// Informational only — the deadline above is what drives the UI.
|
|
||||||
lead = 0
|
|
||||||
}
|
|
||||||
listener.OnSessionExpiring(deadline.Unix(), int64(lead),
|
|
||||||
meta[sessionwatch.MetaSessionFinal] == "true")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) beginExtend() (context.Context, error) {
|
func (c *Client) beginExtend() (context.Context, error) {
|
||||||
@@ -293,11 +222,13 @@ func (c *Client) extendAuthSession(ctx context.Context, urlOpener URLOpener, isA
|
|||||||
}
|
}
|
||||||
defer authClient.Close()
|
defer authClient.Close()
|
||||||
|
|
||||||
// Passing the config path makes the flow pick up the login_hint: an extend
|
// Passing the config path makes the flow pick up the login_hint. That alone
|
||||||
// renews the session of the account already signed in, so it must not stop to
|
// cannot keep the IdP on this profile's account though — a hint is only a
|
||||||
// offer a choice.
|
// 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)
|
a := NewAuthWithConfig(ctx, cfg, cfgPath)
|
||||||
tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV)
|
tokenInfo, err := a.foregroundGetTokenInfoFlow(authClient, urlOpener, isAndroidTV, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("interactive sso login failed: %v", err)
|
return fmt.Errorf("interactive sso login failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -31,6 +31,8 @@ const (
|
|||||||
// PasswordRequiredMarker tells Java to prompt for a password and retry. It is
|
// PasswordRequiredMarker tells Java to prompt for a password and retry. It is
|
||||||
// a string because gomobile flattens errors to their message, so a sentinel
|
// a string because gomobile flattens errors to their message, so a sentinel
|
||||||
// value would not survive the binding.
|
// value would not survive the binding.
|
||||||
|
//
|
||||||
|
//nolint:gosec // G101 false positive: a sentinel marker, not a credential
|
||||||
const PasswordRequiredMarker = "netbird-ssh-password-required"
|
const PasswordRequiredMarker = "netbird-ssh-password-required"
|
||||||
|
|
||||||
// HostKeyUnknownMarker tells Java to show the fingerprint and, on confirmation,
|
// HostKeyUnknownMarker tells Java to show the fingerprint and, on confirmation,
|
||||||
@@ -467,7 +469,7 @@ func (s *SSHClient) requestJWTToken(cfg *profilemanager.Config, cfgPath string)
|
|||||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
|
||||||
defer cancel()
|
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 {
|
if err != nil {
|
||||||
return "", fmt.Errorf("create oauth flow: %w", err)
|
return "", fmt.Errorf("create oauth flow: %w", err)
|
||||||
}
|
}
|
||||||
|
|||||||
+58
-26
@@ -23,7 +23,10 @@ import (
|
|||||||
"github.com/netbirdio/netbird/version"
|
"github.com/netbirdio/netbird/version"
|
||||||
)
|
)
|
||||||
|
|
||||||
const errCloseConnection = "Failed to close connection: %v"
|
const (
|
||||||
|
errCloseConnection = "Failed to close connection: %v"
|
||||||
|
noUpDownFlag = "no-updown"
|
||||||
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
logFileCount uint32
|
logFileCount uint32
|
||||||
@@ -257,13 +260,14 @@ func runForDuration(cmd *cobra.Command, args []string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
stateWasDown := stat.Status != string(internal.StatusConnected) && stat.Status != string(internal.StatusConnecting)
|
stateWasDown := stat.Status != string(internal.StatusConnected) && stat.Status != string(internal.StatusConnecting)
|
||||||
|
noUpDown, _ := cmd.Flags().GetBool(noUpDownFlag)
|
||||||
|
|
||||||
initialLogLevel, err := client.GetLogLevel(cmd.Context(), &proto.GetLogLevelRequest{})
|
initialLogLevel, err := client.GetLogLevel(cmd.Context(), &proto.GetLogLevelRequest{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to get log level: %v", status.Convert(err).Message())
|
return fmt.Errorf("failed to get log level: %v", status.Convert(err).Message())
|
||||||
}
|
}
|
||||||
|
|
||||||
if stateWasDown {
|
if stateWasDown && !noUpDown {
|
||||||
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
|
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
|
||||||
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
|
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
|
||||||
} else {
|
} else {
|
||||||
@@ -284,34 +288,20 @@ func runForDuration(cmd *cobra.Command, args []string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
needsRestoreUp := false
|
needsRestoreUp := false
|
||||||
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
|
if noUpDown {
|
||||||
cmd.PrintErrf("Failed to bring service down: %v\n", status.Convert(err).Message())
|
enableSyncResponsePersistence(cmd, client)
|
||||||
} else {
|
} else {
|
||||||
needsRestoreUp = !stateWasDown
|
needsRestoreUp = restartDaemon(cmd, client, stateWasDown)
|
||||||
cmd.Println("netbird down")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
time.Sleep(1 * time.Second)
|
|
||||||
|
|
||||||
// Enable sync response persistence before bringing the service up
|
|
||||||
if _, err := client.SetSyncResponsePersistence(cmd.Context(), &proto.SetSyncResponsePersistenceRequest{
|
|
||||||
Enabled: true,
|
|
||||||
}); err != nil {
|
|
||||||
cmd.PrintErrf("Failed to enable sync response persistence: %v\n", status.Convert(err).Message())
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
|
|
||||||
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
|
|
||||||
} else {
|
|
||||||
needsRestoreUp = false
|
|
||||||
cmd.Println("netbird up")
|
|
||||||
}
|
|
||||||
|
|
||||||
time.Sleep(3 * time.Second)
|
|
||||||
|
|
||||||
cpuProfilingStarted := false
|
cpuProfilingStarted := false
|
||||||
if _, err := client.StartCPUProfile(cmd.Context(), &proto.StartCPUProfileRequest{}); err != nil {
|
if _, err := client.StartCPUProfile(cmd.Context(), &proto.StartCPUProfileRequest{}); err != nil {
|
||||||
cmd.PrintErrf("Failed to start CPU profiling: %v\n", err)
|
if msg := status.Convert(err).Message(); strings.Contains(msg, "already in progress") {
|
||||||
|
cmd.PrintErrln("CPU profiling is already running (started with `netbird debug cpu start`). " +
|
||||||
|
"It is left running and is included in a bundle created after `netbird debug cpu stop`.")
|
||||||
|
} else {
|
||||||
|
cmd.PrintErrf("Failed to start CPU profiling: %v\n", msg)
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
cpuProfilingStarted = true
|
cpuProfilingStarted = true
|
||||||
defer func() {
|
defer func() {
|
||||||
@@ -401,7 +391,7 @@ func runForDuration(cmd *cobra.Command, args []string) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if stateWasDown {
|
if stateWasDown && !noUpDown {
|
||||||
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
|
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
|
||||||
cmd.PrintErrf("Failed to restore service down state: %v\n", status.Convert(err).Message())
|
cmd.PrintErrf("Failed to restore service down state: %v\n", status.Convert(err).Message())
|
||||||
} else {
|
} else {
|
||||||
@@ -458,6 +448,47 @@ func setSyncResponsePersistence(cmd *cobra.Command, args []string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// enableSyncResponsePersistence asks the daemon to keep the latest sync
|
||||||
|
// response so the bundle carries the network map. With a running daemon only
|
||||||
|
// syncs received after the call are kept.
|
||||||
|
func enableSyncResponsePersistence(cmd *cobra.Command, client proto.DaemonServiceClient) {
|
||||||
|
if _, err := client.SetSyncResponsePersistence(cmd.Context(), &proto.SetSyncResponsePersistenceRequest{
|
||||||
|
Enabled: true,
|
||||||
|
}); err != nil {
|
||||||
|
cmd.PrintErrf("Failed to enable sync response persistence: %v\n", status.Convert(err).Message())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// restartDaemon cycles the daemon down and up with sync response persistence
|
||||||
|
// enabled so the bundle carries the network map. It reports whether the
|
||||||
|
// daemon was left down although it was running before, so the caller can
|
||||||
|
// bring it back up.
|
||||||
|
func restartDaemon(cmd *cobra.Command, client proto.DaemonServiceClient, stateWasDown bool) bool {
|
||||||
|
needsRestoreUp := false
|
||||||
|
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
|
||||||
|
cmd.PrintErrf("Failed to bring service down: %v\n", status.Convert(err).Message())
|
||||||
|
} else {
|
||||||
|
needsRestoreUp = !stateWasDown
|
||||||
|
cmd.Println("netbird down")
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(1 * time.Second)
|
||||||
|
|
||||||
|
// Enable sync response persistence before bringing the service up
|
||||||
|
enableSyncResponsePersistence(cmd, client)
|
||||||
|
|
||||||
|
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
|
||||||
|
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
|
||||||
|
} else {
|
||||||
|
needsRestoreUp = false
|
||||||
|
cmd.Println("netbird up")
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(3 * time.Second)
|
||||||
|
|
||||||
|
return needsRestoreUp
|
||||||
|
}
|
||||||
|
|
||||||
func waitForDurationOrCancel(ctx context.Context, duration time.Duration, cmd *cobra.Command) error {
|
func waitForDurationOrCancel(ctx context.Context, duration time.Duration, cmd *cobra.Command) error {
|
||||||
ticker := time.NewTicker(1 * time.Second)
|
ticker := time.NewTicker(1 * time.Second)
|
||||||
defer ticker.Stop()
|
defer ticker.Stop()
|
||||||
@@ -546,4 +577,5 @@ func init() {
|
|||||||
forCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle")
|
forCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle")
|
||||||
forCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root")
|
forCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root")
|
||||||
forCmd.Flags().Bool("capture", false, "Capture packets during the debug duration and include in bundle")
|
forCmd.Flags().Bool("capture", false, "Capture packets during the debug duration and include in bundle")
|
||||||
|
forCmd.Flags().Bool(noUpDownFlag, false, "Keep the daemon running instead of bringing it down and up before collecting. The bundle only includes the network map if a sync arrives during the run")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -0,0 +1,164 @@
|
|||||||
|
package cmd
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"os/user"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/spf13/cobra"
|
||||||
|
"github.com/spf13/pflag"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
)
|
||||||
|
|
||||||
|
// startDebugTestDaemon starts an in-process daemon with an isolated profile
|
||||||
|
// directory and returns the address the CLI should dial.
|
||||||
|
func startDebugTestDaemon(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
tempDir := t.TempDir()
|
||||||
|
origDefaultProfileDir := profilemanager.DefaultConfigPathDir
|
||||||
|
origActiveProfileStatePath := profilemanager.ActiveProfileStatePath
|
||||||
|
origConfigDirOverride := profilemanager.ConfigDirOverride
|
||||||
|
origDaemonAddr := daemonAddr
|
||||||
|
t.Cleanup(func() {
|
||||||
|
profilemanager.DefaultConfigPathDir = origDefaultProfileDir
|
||||||
|
profilemanager.ActiveProfileStatePath = origActiveProfileStatePath
|
||||||
|
profilemanager.ConfigDirOverride = origConfigDirOverride
|
||||||
|
daemonAddr = origDaemonAddr
|
||||||
|
})
|
||||||
|
|
||||||
|
profilemanager.DefaultConfigPathDir = tempDir
|
||||||
|
profilemanager.ActiveProfileStatePath = tempDir + "/active_profile.json"
|
||||||
|
profilemanager.ConfigDirOverride = tempDir
|
||||||
|
|
||||||
|
currUser, err := user.Current()
|
||||||
|
require.NoError(t, err)
|
||||||
|
sm := profilemanager.ServiceManager{}
|
||||||
|
created, err := sm.AddProfile("test1", currUser.Username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, sm.SetActiveProfileState(&profilemanager.ActiveProfileState{
|
||||||
|
ID: created.ID,
|
||||||
|
Username: currUser.Username,
|
||||||
|
}))
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(internal.CtxInitState(context.Background()))
|
||||||
|
srv, lis := startClientDaemon(t, ctx, "", tempDir+"/config.json")
|
||||||
|
t.Cleanup(func() {
|
||||||
|
cancel()
|
||||||
|
srv.Stop()
|
||||||
|
})
|
||||||
|
|
||||||
|
return "tcp://" + lis.Addr().String()
|
||||||
|
}
|
||||||
|
|
||||||
|
// runDebugCmd runs `netbird debug <args>` against the daemon at addr and
|
||||||
|
// returns everything the command printed.
|
||||||
|
func runDebugCmd(addr string, args ...string) (string, error) {
|
||||||
|
daemonAddr = addr
|
||||||
|
var out bytes.Buffer
|
||||||
|
rootCmd.SetOut(&out)
|
||||||
|
rootCmd.SetErr(&out)
|
||||||
|
rootCmd.SetArgs(append(append([]string{"debug"}, args...), "--daemon-addr", addr, "--log-file", ""))
|
||||||
|
err := rootCmd.Execute()
|
||||||
|
rootCmd.SetOut(nil)
|
||||||
|
rootCmd.SetErr(nil)
|
||||||
|
rootCmd.SetArgs(nil)
|
||||||
|
resetFlags(rootCmd)
|
||||||
|
return out.String(), err
|
||||||
|
}
|
||||||
|
|
||||||
|
// resetFlags puts every flag of the command and its subcommands back to its
|
||||||
|
// default so a value parsed in one run does not leak into the next in-process
|
||||||
|
// execution.
|
||||||
|
func resetFlags(cmd *cobra.Command) {
|
||||||
|
reset := func(f *pflag.Flag) {
|
||||||
|
// Set appends to a slice flag and would parse the "[a,b]" default
|
||||||
|
// text as elements, so slices are replaced instead.
|
||||||
|
if sv, ok := f.Value.(pflag.SliceValue); ok {
|
||||||
|
var def []string
|
||||||
|
if trimmed := strings.Trim(f.DefValue, "[]"); trimmed != "" {
|
||||||
|
def = strings.Split(trimmed, ",")
|
||||||
|
}
|
||||||
|
_ = sv.Replace(def)
|
||||||
|
} else {
|
||||||
|
_ = f.Value.Set(f.DefValue)
|
||||||
|
}
|
||||||
|
f.Changed = false
|
||||||
|
}
|
||||||
|
cmd.Flags().VisitAll(reset)
|
||||||
|
cmd.PersistentFlags().VisitAll(reset)
|
||||||
|
// Commands pin their writers to the buffer of the run that first used
|
||||||
|
// them, so a later run would print into the old buffer.
|
||||||
|
cmd.SetOut(nil)
|
||||||
|
cmd.SetErr(nil)
|
||||||
|
for _, sub := range cmd.Commands() {
|
||||||
|
resetFlags(sub)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestResetFlagsSliceDefault guards against Set("[]") on slice flags, which
|
||||||
|
// stores a literal "[]" element instead of the empty default.
|
||||||
|
func TestResetFlagsSliceDefault(t *testing.T) {
|
||||||
|
cmd := &cobra.Command{Use: "x"}
|
||||||
|
var env, withDefault []string
|
||||||
|
cmd.Flags().StringSliceVar(&env, "env", nil, "")
|
||||||
|
cmd.Flags().StringSliceVar(&withDefault, "names", []string{"a", "b"}, "")
|
||||||
|
require.NoError(t, cmd.Flags().Parse([]string{"--env", "K=V", "--names", "c"}))
|
||||||
|
|
||||||
|
resetFlags(cmd)
|
||||||
|
|
||||||
|
assert.Empty(t, env, "slice flag with no default must reset to empty")
|
||||||
|
assert.Equal(t, []string{"a", "b"}, withDefault, "slice flag must reset to its default")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDebugCPUStartStop(t *testing.T) {
|
||||||
|
addr := startDebugTestDaemon(t)
|
||||||
|
|
||||||
|
run := func(args ...string) error {
|
||||||
|
_, err := runDebugCmd(addr, append([]string{"cpu"}, args...)...)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
require.Error(t, run("stop"), "stop without a running profile must fail")
|
||||||
|
require.NoError(t, run("start"))
|
||||||
|
assert.Error(t, run("start"), "second start must be rejected while profiling")
|
||||||
|
require.NoError(t, run("stop"))
|
||||||
|
assert.Error(t, run("stop"), "second stop must be rejected")
|
||||||
|
assert.NoError(t, run("start"), "profiling can be started again after a stop")
|
||||||
|
assert.NoError(t, run("stop"))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDebugForKeepsRunningCPUProfile covers `debug for` started while a
|
||||||
|
// profile from `debug cpu start` is running: it must say so, leave the
|
||||||
|
// profile alone, and still create the bundle.
|
||||||
|
func TestDebugForKeepsRunningCPUProfile(t *testing.T) {
|
||||||
|
addr := startDebugTestDaemon(t)
|
||||||
|
|
||||||
|
_, err := runDebugCmd(addr, "cpu", "start")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
out, err := runDebugCmd(addr, "for", "1s", "-S=false", "--no-updown")
|
||||||
|
require.NoError(t, err, "output: %s", out)
|
||||||
|
assert.Contains(t, out, "CPU profiling is already running", "the conflict must be explained")
|
||||||
|
assert.NotContains(t, out, "rpc error", "the raw RPC error must not reach the user")
|
||||||
|
assert.Contains(t, out, "Local file:", "the bundle must still be created")
|
||||||
|
|
||||||
|
_, err = runDebugCmd(addr, "cpu", "stop")
|
||||||
|
assert.NoError(t, err, "the profile started by the user must still be running")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDebugForNoUpDown(t *testing.T) {
|
||||||
|
addr := startDebugTestDaemon(t)
|
||||||
|
|
||||||
|
out, err := runDebugCmd(addr, "for", "1s", "-S=false", "--no-updown")
|
||||||
|
require.NoError(t, err, "output: %s", out)
|
||||||
|
assert.NotContains(t, out, "netbird down", "--no-updown must not bring the daemon down")
|
||||||
|
assert.NotContains(t, out, "netbird up", "--no-updown must not bring the daemon up")
|
||||||
|
assert.Contains(t, out, "Local file:", "the bundle must still be created")
|
||||||
|
}
|
||||||
@@ -1,98 +0,0 @@
|
|||||||
package cmd
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"sort"
|
|
||||||
|
|
||||||
"github.com/spf13/cobra"
|
|
||||||
"google.golang.org/grpc/status"
|
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/proto"
|
|
||||||
)
|
|
||||||
|
|
||||||
var forwardingRulesCmd = &cobra.Command{
|
|
||||||
Use: "forwarding",
|
|
||||||
Short: "List forwarding rules",
|
|
||||||
Long: `Commands to list forwarding rules.`,
|
|
||||||
}
|
|
||||||
|
|
||||||
var forwardingRulesListCmd = &cobra.Command{
|
|
||||||
Use: "list",
|
|
||||||
Aliases: []string{"ls"},
|
|
||||||
Short: "List forwarding rules",
|
|
||||||
Example: " netbird forwarding list",
|
|
||||||
Long: "Commands to list forwarding rules.",
|
|
||||||
RunE: listForwardingRules,
|
|
||||||
}
|
|
||||||
|
|
||||||
func listForwardingRules(cmd *cobra.Command, _ []string) error {
|
|
||||||
conn, err := getClient(cmd)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
defer conn.Close()
|
|
||||||
|
|
||||||
client := proto.NewDaemonServiceClient(conn)
|
|
||||||
resp, err := client.ForwardingRules(cmd.Context(), &proto.EmptyRequest{})
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to list network: %v", status.Convert(err).Message())
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(resp.GetRules()) == 0 {
|
|
||||||
cmd.Println("No forwarding rules available.")
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
printForwardingRules(cmd, resp.GetRules())
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func printForwardingRules(cmd *cobra.Command, rules []*proto.ForwardingRule) {
|
|
||||||
cmd.Println("Available forwarding rules:")
|
|
||||||
|
|
||||||
// Sort rules by translated address
|
|
||||||
sort.Slice(rules, func(i, j int) bool {
|
|
||||||
if rules[i].GetTranslatedAddress() != rules[j].GetTranslatedAddress() {
|
|
||||||
return rules[i].GetTranslatedAddress() < rules[j].GetTranslatedAddress()
|
|
||||||
}
|
|
||||||
if rules[i].GetProtocol() != rules[j].GetProtocol() {
|
|
||||||
return rules[i].GetProtocol() < rules[j].GetProtocol()
|
|
||||||
}
|
|
||||||
|
|
||||||
return getFirstPort(rules[i].GetDestinationPort()) < getFirstPort(rules[j].GetDestinationPort())
|
|
||||||
})
|
|
||||||
|
|
||||||
var lastIP string
|
|
||||||
for _, rule := range rules {
|
|
||||||
dPort := portToString(rule.GetDestinationPort())
|
|
||||||
tPort := portToString(rule.GetTranslatedPort())
|
|
||||||
if lastIP != rule.GetTranslatedAddress() {
|
|
||||||
lastIP = rule.GetTranslatedAddress()
|
|
||||||
cmd.Printf("\nTranslated peer: %s\n", rule.GetTranslatedHostname())
|
|
||||||
}
|
|
||||||
|
|
||||||
cmd.Printf(" Local %s/%s to %s:%s\n", rule.GetProtocol(), dPort, rule.GetTranslatedAddress(), tPort)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func getFirstPort(portInfo *proto.PortInfo) int {
|
|
||||||
switch v := portInfo.PortSelection.(type) {
|
|
||||||
case *proto.PortInfo_Port:
|
|
||||||
return int(v.Port)
|
|
||||||
case *proto.PortInfo_Range_:
|
|
||||||
return int(v.Range.GetStart())
|
|
||||||
default:
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func portToString(translatedPort *proto.PortInfo) string {
|
|
||||||
switch v := translatedPort.PortSelection.(type) {
|
|
||||||
case *proto.PortInfo_Port:
|
|
||||||
return fmt.Sprintf("%d", v.Port)
|
|
||||||
case *proto.PortInfo_Range_:
|
|
||||||
return fmt.Sprintf("%d-%d", v.Range.GetStart(), v.Range.GetEnd())
|
|
||||||
default:
|
|
||||||
return "No port specified"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+60
-8
@@ -9,12 +9,11 @@ import (
|
|||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
"golang.org/x/term"
|
"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"
|
||||||
"github.com/netbirdio/netbird/client/internal/auth"
|
"github.com/netbirdio/netbird/client/internal/auth"
|
||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
"github.com/netbirdio/netbird/client/mdm"
|
||||||
nbnet "github.com/netbirdio/netbird/client/net"
|
nbnet "github.com/netbirdio/netbird/client/net"
|
||||||
"github.com/netbirdio/netbird/client/proto"
|
"github.com/netbirdio/netbird/client/proto"
|
||||||
"github.com/netbirdio/netbird/client/server"
|
"github.com/netbirdio/netbird/client/server"
|
||||||
@@ -144,10 +143,7 @@ func doDaemonLogin(ctx context.Context, cmd *cobra.Command, providedSetupKey str
|
|||||||
err = WithBackOff(func() error {
|
err = WithBackOff(func() error {
|
||||||
var backOffErr error
|
var backOffErr error
|
||||||
loginResp, backOffErr = client.Login(ctx, &loginRequest)
|
loginResp, backOffErr = client.Login(ctx, &loginRequest)
|
||||||
if s, ok := gstatus.FromError(backOffErr); ok && (s.Code() == codes.InvalidArgument ||
|
if terminalLoginError(backOffErr) {
|
||||||
s.Code() == codes.PermissionDenied ||
|
|
||||||
s.Code() == codes.NotFound ||
|
|
||||||
s.Code() == codes.Unimplemented) {
|
|
||||||
loginErr = backOffErr
|
loginErr = backOffErr
|
||||||
return nil
|
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 {
|
if err != nil {
|
||||||
return fmt.Errorf("read config file %s: %v", configFilePath, err)
|
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,
|
// Mirror runInForegroundMode: recover residual state (DNS, firewall,
|
||||||
// ssh config, legacy routing) from a previous unclean shutdown and
|
// 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
|
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 {
|
if err != nil {
|
||||||
return nil, err
|
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())
|
flowInfo, err := oAuthFlow.RequestAuthInfo(context.TODO())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("getting a request OAuth flow info failed: %v", err)
|
return nil, fmt.Errorf("getting a request OAuth flow info failed: %v", err)
|
||||||
|
|||||||
+39
-3
@@ -20,6 +20,8 @@ import (
|
|||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
"github.com/spf13/pflag"
|
"github.com/spf13/pflag"
|
||||||
"google.golang.org/grpc"
|
"google.golang.org/grpc"
|
||||||
|
"google.golang.org/grpc/codes"
|
||||||
|
gstatus "google.golang.org/grpc/status"
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/anonymize"
|
"github.com/netbirdio/netbird/client/anonymize"
|
||||||
daddr "github.com/netbirdio/netbird/client/internal/daemonaddr"
|
daddr "github.com/netbirdio/netbird/client/internal/daemonaddr"
|
||||||
@@ -175,7 +177,6 @@ func init() {
|
|||||||
rootCmd.AddCommand(versionCmd)
|
rootCmd.AddCommand(versionCmd)
|
||||||
rootCmd.AddCommand(sshCmd)
|
rootCmd.AddCommand(sshCmd)
|
||||||
rootCmd.AddCommand(networksCMD)
|
rootCmd.AddCommand(networksCMD)
|
||||||
rootCmd.AddCommand(forwardingRulesCmd)
|
|
||||||
rootCmd.AddCommand(debugCmd)
|
rootCmd.AddCommand(debugCmd)
|
||||||
rootCmd.AddCommand(profileCmd)
|
rootCmd.AddCommand(profileCmd)
|
||||||
rootCmd.AddCommand(exposeCmd)
|
rootCmd.AddCommand(exposeCmd)
|
||||||
@@ -183,8 +184,6 @@ func init() {
|
|||||||
networksCMD.AddCommand(routesListCmd)
|
networksCMD.AddCommand(routesListCmd)
|
||||||
networksCMD.AddCommand(routesSelectCmd, routesDeselectCmd)
|
networksCMD.AddCommand(routesSelectCmd, routesDeselectCmd)
|
||||||
|
|
||||||
forwardingRulesCmd.AddCommand(forwardingRulesListCmd)
|
|
||||||
|
|
||||||
debugCmd.AddCommand(debugBundleCmd)
|
debugCmd.AddCommand(debugBundleCmd)
|
||||||
debugCmd.AddCommand(logCmd)
|
debugCmd.AddCommand(logCmd)
|
||||||
logCmd.AddCommand(logLevelCmd)
|
logCmd.AddCommand(logLevelCmd)
|
||||||
@@ -285,6 +284,43 @@ func DialClientGRPCServer(ctx context.Context, addr string) (*grpc.ClientConn, e
|
|||||||
return grpc.DialContext(ctx, target, opts...)
|
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.
|
// WithBackOff execute function in backoff cycle.
|
||||||
func WithBackOff(bf func() error) error {
|
func WithBackOff(bf func() error) error {
|
||||||
return backoff.RetryNotify(bf, CLIBackOffSettings, func(err error, duration time.Duration) {
|
return backoff.RetryNotify(bf, CLIBackOffSettings, func(err error, duration time.Duration) {
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"runtime"
|
"runtime"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
@@ -25,6 +26,30 @@ var serviceCmd = &cobra.Command{
|
|||||||
|
|
||||||
const defaultJSONSocket = "unix:///var/run/netbird-http.sock"
|
const defaultJSONSocket = "unix:///var/run/netbird-http.sock"
|
||||||
|
|
||||||
|
// forbiddenServiceEnvVars are the environment variables the service is never
|
||||||
|
// registered with, keyed in upper case since these are Windows names. Each one
|
||||||
|
// decides where the daemon resolves something it then uses with the privileges
|
||||||
|
// of the account it runs under — LocalSystem on Windows, root elsewhere: the
|
||||||
|
// executables it runs (PATH, PATHEXT, COMSPEC, SystemRoot, windir) or the
|
||||||
|
// directory it writes temporary files in (TEMP, TMP). The daemon needs none of
|
||||||
|
// them, and the utilities it shells out to are resolved by absolute path.
|
||||||
|
var forbiddenServiceEnvVars = map[string]struct{}{
|
||||||
|
"PATH": {},
|
||||||
|
"PATHEXT": {},
|
||||||
|
"SYSTEMROOT": {},
|
||||||
|
"WINDIR": {},
|
||||||
|
"COMSPEC": {},
|
||||||
|
"TEMP": {},
|
||||||
|
"TMP": {},
|
||||||
|
}
|
||||||
|
|
||||||
|
// forbiddenServiceEnvPrefixes are the dynamic-loader families, refused whole
|
||||||
|
// rather than by name: LD_PRELOAD, DYLD_INSERT_LIBRARIES and their siblings all
|
||||||
|
// reach the loader of the process, the set differs per platform and libc, and
|
||||||
|
// new members arrive with new OS releases. Listing them one by one is a list
|
||||||
|
// that is wrong the moment it is written.
|
||||||
|
var forbiddenServiceEnvPrefixes = []string{"LD_", "DYLD_"}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
serviceName string
|
serviceName string
|
||||||
serviceEnvVars []string
|
serviceEnvVars []string
|
||||||
@@ -127,8 +152,33 @@ func parseServiceEnvVars(envVars []string) (map[string]string, error) {
|
|||||||
return nil, fmt.Errorf("empty environment variable key in: %s", env)
|
return nil, fmt.Errorf("empty environment variable key in: %s", env)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if isForbiddenServiceEnvVar(key) {
|
||||||
|
return nil, fmt.Errorf("environment variable %s cannot be set on the service: it decides where the service resolves the executables, libraries or temporary files it uses", key)
|
||||||
|
}
|
||||||
|
|
||||||
envMap[key] = value
|
envMap[key] = value
|
||||||
}
|
}
|
||||||
|
|
||||||
return envMap, nil
|
return envMap, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// isForbiddenServiceEnvVar reports whether name is one the service must not be
|
||||||
|
// registered with.
|
||||||
|
//
|
||||||
|
// The names are matched case-insensitively only on Windows, where they are the
|
||||||
|
// same variable however they are spelled. Elsewhere the environment is
|
||||||
|
// case-sensitive, so Path and PATH are two different variables and only the
|
||||||
|
// exact spelling is the one the loader reads.
|
||||||
|
func isForbiddenServiceEnvVar(name string) bool {
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
name = strings.ToUpper(name)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, forbidden := forbiddenServiceEnvVars[name]; forbidden {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
return slices.ContainsFunc(forbiddenServiceEnvPrefixes, func(prefix string) bool {
|
||||||
|
return strings.HasPrefix(name, prefix)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import (
|
|||||||
|
|
||||||
"github.com/netbirdio/netbird/client/configs"
|
"github.com/netbirdio/netbird/client/configs"
|
||||||
"github.com/netbirdio/netbird/client/internal/daemonaddr"
|
"github.com/netbirdio/netbird/client/internal/daemonaddr"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/elevate"
|
||||||
"github.com/netbirdio/netbird/util"
|
"github.com/netbirdio/netbird/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -43,10 +44,33 @@ func serviceParamsPath() string {
|
|||||||
|
|
||||||
// loadServiceParams reads saved service parameters from disk.
|
// loadServiceParams reads saved service parameters from disk.
|
||||||
// Returns nil with no error if the file does not exist.
|
// Returns nil with no error if the file does not exist.
|
||||||
|
//
|
||||||
|
// The file is read by an elevated install and decides the arguments and the
|
||||||
|
// environment of the service it then registers, so it is used only when its
|
||||||
|
// ownership and permissions are the ones saveServiceParams leaves behind. That
|
||||||
|
// restricted ACL is applied when the file is written, which is not necessarily
|
||||||
|
// before it is first read, so this is checked rather than assumed. A file that
|
||||||
|
// fails the check is treated as absent, and the install proceeds with its
|
||||||
|
// defaults.
|
||||||
func loadServiceParams() (*serviceParams, error) {
|
func loadServiceParams() (*serviceParams, error) {
|
||||||
path := serviceParamsPath()
|
path := serviceParamsPath()
|
||||||
|
|
||||||
data, err := os.ReadFile(path)
|
// Resolve links first so the checks apply to the file that is actually read.
|
||||||
|
// Since the check covers every directory above it as well, nobody who fails
|
||||||
|
// it can swap the file between here and the read below.
|
||||||
|
resolved, err := filepath.EvalSymlinks(path)
|
||||||
|
if err != nil {
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return nil, nil //nolint:nilnil
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("resolve service params %s: %w", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := elevate.CheckOnlyOwnerWritable(resolved); err != nil {
|
||||||
|
return nil, fmt.Errorf("refusing to read service params from %s: %w", resolved, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := os.ReadFile(resolved)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if os.IsNotExist(err) {
|
if os.IsNotExist(err) {
|
||||||
return nil, nil //nolint:nilnil
|
return nil, nil //nolint:nilnil
|
||||||
@@ -182,10 +206,16 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) {
|
|||||||
// If --service-env was explicitly set to empty, all saved env vars are cleared.
|
// If --service-env was explicitly set to empty, all saved env vars are cleared.
|
||||||
// If --service-env was not set, saved env vars are used entirely.
|
// If --service-env was not set, saved env vars are used entirely.
|
||||||
func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
|
func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
|
||||||
|
// A forbidden name explicitly passed on the command line is an error the
|
||||||
|
// operator is told about, but one restored from a file written by an older
|
||||||
|
// version is dropped: an install that refuses to run would leave the host
|
||||||
|
// without a daemon over a variable nobody is asking for any more.
|
||||||
|
saved := dropForbiddenServiceEnvVars(cmd, params.ServiceEnvVars)
|
||||||
|
|
||||||
if !cmd.Flags().Changed("service-env") {
|
if !cmd.Flags().Changed("service-env") {
|
||||||
if len(params.ServiceEnvVars) > 0 {
|
if len(saved) > 0 {
|
||||||
// No explicit env vars: rebuild serviceEnvVars from saved params.
|
// No explicit env vars: rebuild serviceEnvVars from saved params.
|
||||||
serviceEnvVars = envMapToSlice(params.ServiceEnvVars)
|
serviceEnvVars = envMapToSlice(saved)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -204,13 +234,13 @@ func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(params.ServiceEnvVars) == 0 {
|
if len(saved) == 0 {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Merge saved values underneath explicit ones.
|
// Merge saved values underneath explicit ones.
|
||||||
merged := make(map[string]string, len(params.ServiceEnvVars)+len(explicit))
|
merged := make(map[string]string, len(saved)+len(explicit))
|
||||||
maps.Copy(merged, params.ServiceEnvVars)
|
maps.Copy(merged, saved)
|
||||||
maps.Copy(merged, explicit) // explicit wins on conflict
|
maps.Copy(merged, explicit) // explicit wins on conflict
|
||||||
serviceEnvVars = envMapToSlice(merged)
|
serviceEnvVars = envMapToSlice(merged)
|
||||||
}
|
}
|
||||||
@@ -233,6 +263,20 @@ var resetParamsCmd = &cobra.Command{
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// dropForbiddenServiceEnvVars returns the saved entries that may still be
|
||||||
|
// registered on the service, reporting every one it leaves behind.
|
||||||
|
func dropForbiddenServiceEnvVars(cmd *cobra.Command, saved map[string]string) map[string]string {
|
||||||
|
kept := make(map[string]string, len(saved))
|
||||||
|
for key, value := range saved {
|
||||||
|
if isForbiddenServiceEnvVar(key) {
|
||||||
|
cmd.PrintErrf("Warning: ignoring saved service environment variable %s: it decides where the service resolves the executables, libraries or temporary files it uses\n", key)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
kept[key] = value
|
||||||
|
}
|
||||||
|
return kept
|
||||||
|
}
|
||||||
|
|
||||||
// envMapToSlice converts a map of env vars to a KEY=VALUE slice.
|
// envMapToSlice converts a map of env vars to a KEY=VALUE slice.
|
||||||
func envMapToSlice(m map[string]string) []string {
|
func envMapToSlice(m map[string]string) []string {
|
||||||
s := make([]string, 0, len(m))
|
s := make([]string, 0, len(m))
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
"go/token"
|
"go/token"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -353,6 +354,59 @@ func TestApplyServiceEnvParams_NotChanged(t *testing.T) {
|
|||||||
assert.Equal(t, map[string]string{"FROM_SAVED": "val"}, result)
|
assert.Equal(t, map[string]string{"FROM_SAVED": "val"}, result)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestParseServiceEnvVars_RejectsForbiddenNames(t *testing.T) {
|
||||||
|
for _, env := range []string{"PATH=C:\\somewhere", "LD_PRELOAD=/tmp/lib.so", "DYLD_FALLBACK_LIBRARY_PATH=/tmp"} {
|
||||||
|
_, err := parseServiceEnvVars([]string{"KEEP=me", env})
|
||||||
|
require.Errorf(t, err, "%s selects what the service resolves and must be refused", env)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsForbiddenServiceEnvVar(t *testing.T) {
|
||||||
|
// The loader families are matched by prefix, so a name nobody has heard of
|
||||||
|
// yet is refused too.
|
||||||
|
for _, name := range []string{
|
||||||
|
"PATH", "PATHEXT", "COMSPEC", "SYSTEMROOT", "WINDIR", "TEMP", "TMP",
|
||||||
|
"LD_PRELOAD", "LD_AUDIT", "DYLD_INSERT_LIBRARIES", "DYLD_FALLBACK_FRAMEWORK_PATH",
|
||||||
|
} {
|
||||||
|
assert.Truef(t, isForbiddenServiceEnvVar(name), "%s must be refused", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The prefix must not swallow names that merely start with the same letters.
|
||||||
|
for _, name := range []string{"NB_LOG_LEVEL", "NB_WG_DEBUG", "HTTPS_PROXY", "LDAP_URL", "DYLDX"} {
|
||||||
|
assert.Falsef(t, isForbiddenServiceEnvVar(name), "%s has no reason to be refused", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// On Windows a variable is the same one however it is spelled; elsewhere
|
||||||
|
// Path and PATH are two variables and only the exact one is read.
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
assert.True(t, isForbiddenServiceEnvVar("Path"))
|
||||||
|
assert.True(t, isForbiddenServiceEnvVar("ld_preload"))
|
||||||
|
} else {
|
||||||
|
assert.False(t, isForbiddenServiceEnvVar("Path"))
|
||||||
|
assert.False(t, isForbiddenServiceEnvVar("ld_preload"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestApplyServiceEnvParams_DropsForbiddenSavedNames(t *testing.T) {
|
||||||
|
origServiceEnvVars := serviceEnvVars
|
||||||
|
t.Cleanup(func() { serviceEnvVars = origServiceEnvVars })
|
||||||
|
|
||||||
|
serviceEnvVars = nil
|
||||||
|
|
||||||
|
cmd := &cobra.Command{}
|
||||||
|
cmd.Flags().StringSlice("service-env", nil, "")
|
||||||
|
|
||||||
|
saved := &serviceParams{
|
||||||
|
ServiceEnvVars: map[string]string{"PATH": "C:\\attacker", "NB_LOG_FORMAT": "json"},
|
||||||
|
}
|
||||||
|
|
||||||
|
applyServiceEnvParams(cmd, saved)
|
||||||
|
|
||||||
|
result, err := parseServiceEnvVars(serviceEnvVars)
|
||||||
|
require.NoError(t, err, "a saved PATH must be dropped rather than fail the install")
|
||||||
|
assert.Equal(t, map[string]string{"NB_LOG_FORMAT": "json"}, result)
|
||||||
|
}
|
||||||
|
|
||||||
func TestApplyServiceEnvParams_ExplicitEmptyClears(t *testing.T) {
|
func TestApplyServiceEnvParams_ExplicitEmptyClears(t *testing.T) {
|
||||||
origServiceEnvVars := serviceEnvVars
|
origServiceEnvVars := serviceEnvVars
|
||||||
t.Cleanup(func() { serviceEnvVars = origServiceEnvVars })
|
t.Cleanup(func() { serviceEnvVars = origServiceEnvVars })
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -6,9 +6,9 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"go.uber.org/mock/gomock"
|
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"go.opentelemetry.io/otel"
|
"go.opentelemetry.io/otel"
|
||||||
|
"go.uber.org/mock/gomock"
|
||||||
"google.golang.org/grpc"
|
"google.golang.org/grpc"
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator"
|
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator"
|
||||||
@@ -28,7 +28,6 @@ import (
|
|||||||
mgmt "github.com/netbirdio/netbird/management/server"
|
mgmt "github.com/netbirdio/netbird/management/server"
|
||||||
"github.com/netbirdio/netbird/management/server/activity"
|
"github.com/netbirdio/netbird/management/server/activity"
|
||||||
"github.com/netbirdio/netbird/management/server/groups"
|
"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/permissions"
|
||||||
"github.com/netbirdio/netbird/management/server/settings"
|
"github.com/netbirdio/netbird/management/server/settings"
|
||||||
"github.com/netbirdio/netbird/management/server/store"
|
"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)
|
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||||
requestBuffer := mgmt.NewAccountRequestBuffer(ctx, store)
|
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 {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|||||||
+33
-7
@@ -21,6 +21,7 @@ import (
|
|||||||
"github.com/netbirdio/netbird/client/internal"
|
"github.com/netbirdio/netbird/client/internal"
|
||||||
"github.com/netbirdio/netbird/client/internal/peer"
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
"github.com/netbirdio/netbird/client/mdm"
|
||||||
nbnet "github.com/netbirdio/netbird/client/net"
|
nbnet "github.com/netbirdio/netbird/client/net"
|
||||||
"github.com/netbirdio/netbird/client/proto"
|
"github.com/netbirdio/netbird/client/proto"
|
||||||
"github.com/netbirdio/netbird/client/server"
|
"github.com/netbirdio/netbird/client/server"
|
||||||
@@ -234,6 +235,10 @@ func runInForegroundMode(ctx context.Context, cmd *cobra.Command, activeProf *pr
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get config file: %v", err)
|
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)
|
_, _ = profilemanager.UpdateOldManagementURL(ctx, config, configFilePath)
|
||||||
|
|
||||||
@@ -352,9 +357,17 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager
|
|||||||
// set the new config
|
// set the new config
|
||||||
req := setupSetConfigReq(customDNSAddressConverted, cmd, activeProf.ID.String(), username.Username)
|
req := setupSetConfigReq(customDNSAddressConverted, cmd, activeProf.ID.String(), username.Username)
|
||||||
if _, err := client.SetConfig(ctx, req); err != nil {
|
if _, err := client.SetConfig(ctx, req); err != nil {
|
||||||
if st, ok := gstatus.FromError(err); ok && st.Code() == codes.Unavailable {
|
switch reason, refused := refusedSettingsUpdate(err); {
|
||||||
log.Warnf("setConfig method is not available in the daemon: %s", st.Message())
|
case refused:
|
||||||
} else {
|
// 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)
|
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 {
|
err = WithBackOff(func() error {
|
||||||
var backOffErr error
|
var backOffErr error
|
||||||
loginResp, backOffErr = client.Login(ctx, loginRequest)
|
loginResp, backOffErr = client.Login(ctx, loginRequest)
|
||||||
if s, ok := gstatus.FromError(backOffErr); ok && (s.Code() == codes.InvalidArgument ||
|
if terminalLoginError(backOffErr) {
|
||||||
s.Code() == codes.PermissionDenied ||
|
|
||||||
s.Code() == codes.NotFound ||
|
|
||||||
s.Code() == codes.Unimplemented) {
|
|
||||||
loginErr = backOffErr
|
loginErr = backOffErr
|
||||||
return nil
|
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 {
|
func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, profileName, username string) *proto.SetConfigRequest {
|
||||||
var req proto.SetConfigRequest
|
var req proto.SetConfigRequest
|
||||||
req.ProfileName = profileName
|
req.ProfileName = profileName
|
||||||
|
|||||||
@@ -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))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -21,6 +21,7 @@ import (
|
|||||||
"github.com/netbirdio/netbird/client/internal/auth"
|
"github.com/netbirdio/netbird/client/internal/auth"
|
||||||
"github.com/netbirdio/netbird/client/internal/peer"
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
"github.com/netbirdio/netbird/client/mdm"
|
||||||
nbssh "github.com/netbirdio/netbird/client/ssh"
|
nbssh "github.com/netbirdio/netbird/client/ssh"
|
||||||
"github.com/netbirdio/netbird/client/system"
|
"github.com/netbirdio/netbird/client/system"
|
||||||
"github.com/netbirdio/netbird/shared/management/domain"
|
"github.com/netbirdio/netbird/shared/management/domain"
|
||||||
@@ -229,6 +230,10 @@ func New(opts Options) (*Client, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("create config: %w", err)
|
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 != "" {
|
if opts.PrivateKey != "" {
|
||||||
config.PrivateKey = opts.PrivateKey
|
config.PrivateKey = opts.PrivateKey
|
||||||
|
|||||||
@@ -6,8 +6,8 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"go.uber.org/mock/gomock"
|
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
"go.uber.org/mock/gomock"
|
||||||
"google.golang.org/grpc"
|
"google.golang.org/grpc"
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller"
|
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller"
|
||||||
@@ -21,7 +21,6 @@ import (
|
|||||||
nbcache "github.com/netbirdio/netbird/management/server/cache"
|
nbcache "github.com/netbirdio/netbird/management/server/cache"
|
||||||
"github.com/netbirdio/netbird/management/server/groups"
|
"github.com/netbirdio/netbird/management/server/groups"
|
||||||
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator"
|
"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/job"
|
||||||
"github.com/netbirdio/netbird/management/server/permissions"
|
"github.com/netbirdio/netbird/management/server/permissions"
|
||||||
"github.com/netbirdio/netbird/management/server/settings"
|
"github.com/netbirdio/netbird/management/server/settings"
|
||||||
@@ -146,8 +145,8 @@ func startManagement(t *testing.T, signalAddr string) string {
|
|||||||
|
|
||||||
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
updateManager := update_channel.NewPeersUpdateManager(metrics)
|
||||||
requestBuffer := mgmt.NewAccountRequestBuffer(context.Background(), testStore)
|
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)
|
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, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
|
accountManager, err := mgmt.BuildManager(context.Background(), cfg, testStore, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, settingsMockManager, permissionsManager, false, cacheStore)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
secretsManager, err := nbgrpc.NewTimeBasedAuthSecretsManager(updateManager, cfg.TURNConfig, cfg.Relay, settingsMockManager, groupsManager)
|
secretsManager, err := nbgrpc.NewTimeBasedAuthSecretsManager(updateManager, cfg.TURNConfig, cfg.Relay, settingsMockManager, groupsManager)
|
||||||
|
|||||||
@@ -8,177 +8,11 @@ import (
|
|||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/hashicorp/go-multierror"
|
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
nberrors "github.com/netbirdio/netbird/client/errors"
|
|
||||||
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
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 {
|
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))
|
ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
||||||
|
|
||||||
|
|||||||
@@ -1,240 +0,0 @@
|
|||||||
//go:build privileged
|
|
||||||
|
|
||||||
package iptables
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
fw "github.com/netbirdio/netbird/client/firewall/manager"
|
|
||||||
"github.com/netbirdio/netbird/client/iface"
|
|
||||||
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
|
||||||
)
|
|
||||||
|
|
||||||
func iptRefcountIfaceV4() *iFaceMock {
|
|
||||||
return &iFaceMock{
|
|
||||||
NameFunc: func() string { return "wt-refcount" },
|
|
||||||
AddressFunc: func() wgaddr.Address {
|
|
||||||
return wgaddr.Address{
|
|
||||||
IP: netip.MustParseAddr("10.20.0.1"),
|
|
||||||
Network: netip.MustParsePrefix("10.20.0.0/24"),
|
|
||||||
}
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func iptRefcountIfaceDual() *iFaceMock {
|
|
||||||
return &iFaceMock{
|
|
||||||
NameFunc: func() string { return "wt-refcount" },
|
|
||||||
AddressFunc: func() wgaddr.Address {
|
|
||||||
return wgaddr.Address{
|
|
||||||
IP: netip.MustParseAddr("10.20.0.1"),
|
|
||||||
Network: netip.MustParsePrefix("10.20.0.0/24"),
|
|
||||||
IPv6: netip.MustParseAddr("fd00::1"),
|
|
||||||
IPv6Net: netip.MustParsePrefix("fd00::/64"),
|
|
||||||
}
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newIptRefcountManager(t *testing.T, dual bool) *Manager {
|
|
||||||
t.Helper()
|
|
||||||
var ifMock *iFaceMock
|
|
||||||
if dual {
|
|
||||||
ifMock = iptRefcountIfaceDual()
|
|
||||||
} else {
|
|
||||||
ifMock = iptRefcountIfaceV4()
|
|
||||||
}
|
|
||||||
m, err := Create(ifMock, iface.DefaultMTU)
|
|
||||||
require.NoError(t, err, "create manager")
|
|
||||||
require.NoError(t, m.Init(nil), "init manager")
|
|
||||||
t.Cleanup(func() {
|
|
||||||
require.NoError(t, m.Close(nil), "close manager")
|
|
||||||
})
|
|
||||||
return m
|
|
||||||
}
|
|
||||||
|
|
||||||
func iptDnatV4(port uint16) fw.ForwardRule {
|
|
||||||
return fw.ForwardRule{
|
|
||||||
Protocol: fw.ProtocolTCP,
|
|
||||||
DestinationPort: fw.Port{Values: []uint16{port}},
|
|
||||||
TranslatedAddress: netip.MustParseAddr("10.20.0.2"),
|
|
||||||
TranslatedPort: fw.Port{Values: []uint16{80}},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func iptDnatV6(port uint16) fw.ForwardRule {
|
|
||||||
return fw.ForwardRule{
|
|
||||||
Protocol: fw.ProtocolTCP,
|
|
||||||
DestinationPort: fw.Port{Values: []uint16{port}},
|
|
||||||
TranslatedAddress: netip.MustParseAddr("fd00::2"),
|
|
||||||
TranslatedPort: fw.Port{Values: []uint16{80}},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestIptablesRouting_RepeatedEnableSingleReference verifies that EnableRouting
|
|
||||||
// (called on every network-map update) holds at most one reference per family
|
|
||||||
// and a single DisableRouting drops both back to zero.
|
|
||||||
func TestIptablesRouting_RepeatedEnableSingleReference(t *testing.T) {
|
|
||||||
m := newIptRefcountManager(t, true)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
require.NoError(t, m.EnableRouting(), "first enable")
|
|
||||||
require.NoError(t, m.EnableRouting(), "second enable")
|
|
||||||
require.NoError(t, m.EnableRouting(), "third enable")
|
|
||||||
v4, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 1, v4, "repeated enable holds a single v4 reference")
|
|
||||||
assert.Equal(t, 1, v6, "repeated enable holds a single v6 reference")
|
|
||||||
|
|
||||||
require.NoError(t, m.DisableRouting(), "disable")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4, "single disable releases the v4 reference")
|
|
||||||
assert.Equal(t, 0, v6, "single disable releases the v6 reference")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestIptablesRouting_DisableKeepsDNATReference verifies that an unpaired
|
|
||||||
// DisableRouting does not release references held by active DNAT rules.
|
|
||||||
func TestIptablesRouting_DisableKeepsDNATReference(t *testing.T) {
|
|
||||||
m := newIptRefcountManager(t, true)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(iptDnatV6(9095))
|
|
||||||
require.NoError(t, err, "add v6 dnat")
|
|
||||||
|
|
||||||
require.NoError(t, m.DisableRouting(), "unpaired disable")
|
|
||||||
_, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat")
|
|
||||||
_, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v6, "delete releases the DNAT reference")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestIptablesDNAT_RefcountBalancedV4 covers a Balanced Add/Delete pair on v4.
|
|
||||||
func TestIptablesDNAT_RefcountBalancedV4(t *testing.T) {
|
|
||||||
m := newIptRefcountManager(t, false)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(iptDnatV4(7081))
|
|
||||||
require.NoError(t, err, "add v4 dnat 1")
|
|
||||||
v4, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 1, v4, "v4 refcount after first add")
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount unchanged")
|
|
||||||
|
|
||||||
r2, err := m.AddDNATRule(iptDnatV4(7082))
|
|
||||||
require.NoError(t, err, "add v4 dnat 2")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 2, v4, "v4 refcount after second add")
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount unchanged")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1))
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 1, v4, "v4 refcount after first delete")
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount unchanged")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r2))
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4, "v4 refcount after second delete")
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount unchanged")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestIptablesDNAT_RefcountBalancedV6 checks the v6 path increments v6 only and
|
|
||||||
// decrements back to zero.
|
|
||||||
func TestIptablesDNAT_RefcountBalancedV6(t *testing.T) {
|
|
||||||
m := newIptRefcountManager(t, true)
|
|
||||||
require.NotNil(t, m.family6, "v6 family")
|
|
||||||
require.Same(t, m.family4.ipFwdState, m.family6.ipFwdState, "shared state")
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(iptDnatV6(9081))
|
|
||||||
require.NoError(t, err, "add v6 dnat 1")
|
|
||||||
v4, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 0, v4)
|
|
||||||
assert.Equal(t, 1, v6, "v6 refcount after first add")
|
|
||||||
|
|
||||||
r2, err := m.AddDNATRule(iptDnatV6(9082))
|
|
||||||
require.NoError(t, err, "add v6 dnat 2")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4, "v4 refcount unchanged")
|
|
||||||
assert.Equal(t, 2, v6, "v6 refcount after second add")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1))
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4, "v4 refcount unchanged")
|
|
||||||
assert.Equal(t, 1, v6, "v6 refcount after first delete")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r2))
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4)
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount after second delete")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestIptablesDNAT_DuplicateAddNoLeak verifies the duplicate-rule path returns
|
|
||||||
// without bumping the refcount.
|
|
||||||
func TestIptablesDNAT_DuplicateAddNoLeak(t *testing.T) {
|
|
||||||
m := newIptRefcountManager(t, true)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
rule := iptDnatV4(7083)
|
|
||||||
r1, err := m.AddDNATRule(rule)
|
|
||||||
require.NoError(t, err)
|
|
||||||
v4, _ := state.Counts()
|
|
||||||
assert.Equal(t, 1, v4)
|
|
||||||
|
|
||||||
_, err = m.AddDNATRule(rule)
|
|
||||||
require.NoError(t, err, "duplicate add")
|
|
||||||
v4, _ = state.Counts()
|
|
||||||
assert.Equal(t, 1, v4, "duplicate add must not increment")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1))
|
|
||||||
v4, _ = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4, "single delete must drop to zero")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestIptablesDNAT_DeleteMissingNoUnderflow verifies Delete on an unknown rule
|
|
||||||
// neither errors nor releases the refcount.
|
|
||||||
func TestIptablesDNAT_DeleteMissingNoUnderflow(t *testing.T) {
|
|
||||||
m := newIptRefcountManager(t, true)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
phantom := iptDnatV4(7099)
|
|
||||||
require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4")
|
|
||||||
v4, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 0, v4)
|
|
||||||
assert.Equal(t, 0, v6)
|
|
||||||
|
|
||||||
phantom6 := iptDnatV6(9099)
|
|
||||||
require.NoError(t, m.DeleteDNATRule(&phantom6), "delete missing v6")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4)
|
|
||||||
assert.Equal(t, 0, v6)
|
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(iptDnatV4(7100))
|
|
||||||
require.NoError(t, err)
|
|
||||||
v4, _ = state.Counts()
|
|
||||||
assert.Equal(t, 1, v4, "real add still increments after phantom delete")
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1))
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestIptablesDNAT_DoubleDeleteNoUnderflow verifies a second Delete on the same
|
|
||||||
// rule is a no-op.
|
|
||||||
func TestIptablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) {
|
|
||||||
m := newIptRefcountManager(t, true)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(iptDnatV6(9083))
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 1, v6)
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1), "first delete")
|
|
||||||
_, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v6)
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1), "second delete must be no-op")
|
|
||||||
_, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v6, "double delete must not underflow")
|
|
||||||
}
|
|
||||||
@@ -24,6 +24,7 @@ const (
|
|||||||
tableFilter = "filter"
|
tableFilter = "filter"
|
||||||
tableNat = "nat"
|
tableNat = "nat"
|
||||||
tableMangle = "mangle"
|
tableMangle = "mangle"
|
||||||
|
tableRaw = "raw"
|
||||||
|
|
||||||
// chainACLInput is the peer ACL chain that holds installed
|
// chainACLInput is the peer ACL chain that holds installed
|
||||||
// peer-filtering rules.
|
// peer-filtering rules.
|
||||||
@@ -34,6 +35,7 @@ const (
|
|||||||
mangleForwardKey chainKey = "MANGLE-FORWARD"
|
mangleForwardKey chainKey = "MANGLE-FORWARD"
|
||||||
|
|
||||||
chainInput = "INPUT"
|
chainInput = "INPUT"
|
||||||
|
chainOutput = "OUTPUT"
|
||||||
chainPostrouting = "POSTROUTING"
|
chainPostrouting = "POSTROUTING"
|
||||||
chainPrerouting = "PREROUTING"
|
chainPrerouting = "PREROUTING"
|
||||||
chainForward = "FORWARD"
|
chainForward = "FORWARD"
|
||||||
@@ -54,10 +56,6 @@ const (
|
|||||||
markManglePost = "mark-mangle-post"
|
markManglePost = "mark-mangle-post"
|
||||||
matchSet = "--match-set"
|
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 is the minimum IPv4 (20) + TCP (20) header size for MSS calculation.
|
||||||
ipv4TCPHeaderSize = 40
|
ipv4TCPHeaderSize = 40
|
||||||
// ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation.
|
// ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation.
|
||||||
|
|||||||
@@ -81,15 +81,6 @@ func (r *family) hasRule(id nbid.RuleID) bool {
|
|||||||
return ok
|
return ok
|
||||||
}
|
}
|
||||||
|
|
||||||
// hasDNATRule reports whether this family owns the DNAT rule set for
|
|
||||||
// the given user id. DNAT rules live in r.rules under the well-known
|
|
||||||
// "<id>_dnat" key; the lookup here is used by Manager.DeleteDNATRule
|
|
||||||
// to pick the right family.
|
|
||||||
func (r *family) hasDNATRule(id firewall.RuleID) bool {
|
|
||||||
_, ok := r.rules[id+dnatSuffix]
|
|
||||||
return ok
|
|
||||||
}
|
|
||||||
|
|
||||||
// DeleteFilterRule removes a previously installed filter rule. The
|
// DeleteFilterRule removes a previously installed filter rule. The
|
||||||
// rule's stored chain/table identify where to delete from; source set
|
// rule's stored chain/table identify where to delete from; source set
|
||||||
// references are recovered from the spec via findSets and dropped
|
// references are recovered from the spec via findSets and dropped
|
||||||
|
|||||||
@@ -27,7 +27,6 @@ type Manager struct {
|
|||||||
|
|
||||||
ipv4Client *iptables.IPTables
|
ipv4Client *iptables.IPTables
|
||||||
family4 *family
|
family4 *family
|
||||||
rawSupported bool
|
|
||||||
|
|
||||||
// IPv6 counterparts, nil when no v6 overlay
|
// IPv6 counterparts, nil when no v6 overlay
|
||||||
ipv6Client *iptables.IPTables
|
ipv6Client *iptables.IPTables
|
||||||
@@ -108,10 +107,6 @@ func (m *Manager) Init(stateManager *statemanager.Manager) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := m.initNoTrackChain(); err != nil {
|
|
||||||
log.Warnf("raw table not available, notrack rules will be disabled: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Trust after all fatal init steps so a later failure doesn't leave the
|
// Trust after all fatal init steps so a later failure doesn't leave the
|
||||||
// interface in firewalld's trusted zone without a corresponding Close.
|
// interface in firewalld's trusted zone without a corresponding Close.
|
||||||
if err := firewalld.TrustInterface(m.wgIface.Name()); err != nil {
|
if err := firewalld.TrustInterface(m.wgIface.Name()); err != nil {
|
||||||
@@ -285,10 +280,6 @@ func (m *Manager) Close(stateManager *statemanager.Manager) error {
|
|||||||
|
|
||||||
var merr *multierror.Error
|
var merr *multierror.Error
|
||||||
|
|
||||||
if err := m.cleanupNoTrackChain(); err != nil {
|
|
||||||
merr = multierror.Append(merr, fmt.Errorf("cleanup notrack chain: %w", err))
|
|
||||||
}
|
|
||||||
|
|
||||||
if m.hasIPv6() {
|
if m.hasIPv6() {
|
||||||
if err := m.family6.Reset(); err != nil {
|
if err := m.family6.Reset(); err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("reset v6 family: %w", err))
|
merr = multierror.Append(merr, fmt.Errorf("reset v6 family: %w", err))
|
||||||
@@ -332,31 +323,6 @@ func (m *Manager) DisableRouting() error {
|
|||||||
return m.family4.ipFwdState.ReleaseRouting()
|
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
|
// UpdateSet updates the set with the given prefixes
|
||||||
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
|
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
|
||||||
m.mutex.Lock()
|
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)
|
return m.family4.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
|
||||||
chainNameRaw = "NETBIRD-RAW"
|
|
||||||
chainOutput = "OUTPUT"
|
|
||||||
tableRaw = "raw"
|
|
||||||
)
|
|
||||||
|
|
||||||
// SetupEBPFProxyNoTrack creates notrack rules for eBPF proxy loopback traffic.
|
|
||||||
// This prevents conntrack from tracking WireGuard proxy traffic on loopback, which
|
|
||||||
// can interfere with MASQUERADE rules (e.g., from container runtimes like Podman/netavark).
|
|
||||||
//
|
|
||||||
// Traffic flows that need NOTRACK:
|
|
||||||
//
|
|
||||||
// 1. Egress: WireGuard -> fake endpoint (before eBPF rewrite)
|
|
||||||
// src=127.0.0.1:wgPort -> dst=127.0.0.1:fakePort
|
|
||||||
// Matched by: sport=wgPort
|
|
||||||
//
|
|
||||||
// 2. Egress: Proxy -> WireGuard (via raw socket)
|
|
||||||
// src=127.0.0.1:fakePort -> dst=127.0.0.1:wgPort
|
|
||||||
// Matched by: dport=wgPort
|
|
||||||
//
|
|
||||||
// 3. Ingress: Packets to WireGuard
|
|
||||||
// dst=127.0.0.1:wgPort
|
|
||||||
// Matched by: dport=wgPort
|
|
||||||
//
|
|
||||||
// 4. Ingress: Packets to proxy (after eBPF rewrite)
|
|
||||||
// dst=127.0.0.1:proxyPort
|
|
||||||
// Matched by: dport=proxyPort
|
|
||||||
//
|
|
||||||
// Rules are cleaned up when the firewall manager is closed.
|
|
||||||
func (m *Manager) SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error {
|
|
||||||
m.mutex.Lock()
|
|
||||||
defer m.mutex.Unlock()
|
|
||||||
|
|
||||||
if !m.rawSupported {
|
|
||||||
return fmt.Errorf("raw table not available")
|
|
||||||
}
|
|
||||||
|
|
||||||
wgPortStr := fmt.Sprintf("%d", wgPort)
|
|
||||||
proxyPortStr := fmt.Sprintf("%d", proxyPort)
|
|
||||||
|
|
||||||
// Egress rules: match outgoing loopback UDP packets
|
|
||||||
outputRuleSport := []string{"-o", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--sport", wgPortStr, "-j", "NOTRACK"}
|
|
||||||
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, outputRuleSport...); err != nil {
|
|
||||||
return fmt.Errorf("add output sport notrack rule: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
outputRuleDport := []string{"-o", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", wgPortStr, "-j", "NOTRACK"}
|
|
||||||
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, outputRuleDport...); err != nil {
|
|
||||||
return fmt.Errorf("add output dport notrack rule: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Ingress rules: match incoming loopback UDP packets
|
|
||||||
preroutingRuleWg := []string{"-i", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", wgPortStr, "-j", "NOTRACK"}
|
|
||||||
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, preroutingRuleWg...); err != nil {
|
|
||||||
return fmt.Errorf("add prerouting wg notrack rule: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
preroutingRuleProxy := []string{"-i", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", proxyPortStr, "-j", "NOTRACK"}
|
|
||||||
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, preroutingRuleProxy...); err != nil {
|
|
||||||
return fmt.Errorf("add prerouting proxy notrack rule: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Debugf("set up ebpf proxy notrack rules for ports %d,%d", proxyPort, wgPort)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Manager) initNoTrackChain() error {
|
|
||||||
if err := m.cleanupNoTrackChain(); err != nil {
|
|
||||||
log.Debugf("cleanup notrack chain: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.ipv4Client.NewChain(tableRaw, chainNameRaw); err != nil {
|
|
||||||
return fmt.Errorf("create chain: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
jumpRule := []string{"-j", chainNameRaw}
|
|
||||||
|
|
||||||
if err := m.ipv4Client.InsertUnique(tableRaw, chainOutput, 1, jumpRule...); err != nil {
|
|
||||||
if delErr := m.ipv4Client.DeleteChain(tableRaw, chainNameRaw); delErr != nil {
|
|
||||||
log.Debugf("delete orphan chain: %v", delErr)
|
|
||||||
}
|
|
||||||
return fmt.Errorf("add output jump rule: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.ipv4Client.InsertUnique(tableRaw, chainPrerouting, 1, jumpRule...); err != nil {
|
|
||||||
if delErr := m.ipv4Client.DeleteIfExists(tableRaw, chainOutput, jumpRule...); delErr != nil {
|
|
||||||
log.Debugf("delete output jump rule: %v", delErr)
|
|
||||||
}
|
|
||||||
if delErr := m.ipv4Client.DeleteChain(tableRaw, chainNameRaw); delErr != nil {
|
|
||||||
log.Debugf("delete orphan chain: %v", delErr)
|
|
||||||
}
|
|
||||||
return fmt.Errorf("add prerouting jump rule: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.rawSupported = true
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Manager) cleanupNoTrackChain() error {
|
|
||||||
exists, err := m.ipv4Client.ChainExists(tableRaw, chainNameRaw)
|
|
||||||
if err != nil {
|
|
||||||
if !m.rawSupported {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return fmt.Errorf("check chain exists: %w", err)
|
|
||||||
}
|
|
||||||
if !exists {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
jumpRule := []string{"-j", chainNameRaw}
|
|
||||||
|
|
||||||
if err := m.ipv4Client.DeleteIfExists(tableRaw, chainOutput, jumpRule...); err != nil {
|
|
||||||
return fmt.Errorf("remove output jump rule: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.ipv4Client.DeleteIfExists(tableRaw, chainPrerouting, jumpRule...); err != nil {
|
|
||||||
return fmt.Errorf("remove prerouting jump rule: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.ipv4Client.ClearAndDeleteChain(tableRaw, chainNameRaw); err != nil {
|
|
||||||
return fmt.Errorf("clear and delete chain: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.rawSupported = false
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func getConntrackEstablished() []string {
|
func getConntrackEstablished() []string {
|
||||||
return []string{"-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"}
|
return []string{"-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -497,16 +497,6 @@ func TestIptablesCloseRemovesAllState(t *testing.T) {
|
|||||||
require.NoError(t, manager.AddNatRule(pair), "add nat rule")
|
require.NoError(t, manager.AddNatRule(pair), "add nat rule")
|
||||||
require.NoError(t, manager.EnableRouting(), "enable routing")
|
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")
|
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.
|
// Everything above stays in place, so Close is what has to remove it.
|
||||||
|
|||||||
@@ -172,12 +172,6 @@ type Manager interface {
|
|||||||
|
|
||||||
DisableRouting() error
|
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 updates the set with the given prefixes
|
||||||
UpdateSet(hash Set, prefixes []netip.Prefix) error
|
UpdateSet(hash Set, prefixes []netip.Prefix) error
|
||||||
|
|
||||||
@@ -192,10 +186,6 @@ type Manager interface {
|
|||||||
|
|
||||||
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
|
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
|
||||||
RemoveOutputDNAT(localAddr netip.Addr, protocol Protocol, originalPort, translatedPort uint16) error
|
RemoveOutputDNAT(localAddr netip.Addr, protocol Protocol, originalPort, translatedPort uint16) error
|
||||||
|
|
||||||
// SetupEBPFProxyNoTrack creates static notrack rules for eBPF proxy loopback traffic.
|
|
||||||
// This prevents conntrack from interfering with WireGuard proxy communication.
|
|
||||||
SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// GenKey builds the rule id for this pair from the given format.
|
// GenKey builds the rule id for this pair from the given format.
|
||||||
|
|||||||
@@ -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())
|
|
||||||
}
|
|
||||||
@@ -9,332 +9,11 @@ import (
|
|||||||
"github.com/google/nftables"
|
"github.com/google/nftables"
|
||||||
"github.com/google/nftables/binaryutil"
|
"github.com/google/nftables/binaryutil"
|
||||||
"github.com/google/nftables/expr"
|
"github.com/google/nftables/expr"
|
||||||
"github.com/google/nftables/xt"
|
|
||||||
"github.com/hashicorp/go-multierror"
|
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
nberrors "github.com/netbirdio/netbird/client/errors"
|
|
||||||
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (r *family) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) {
|
|
||||||
ruleID := rule.ID()
|
|
||||||
if _, exists := r.rules[ruleID+dnatSuffix]; exists {
|
|
||||||
return rule, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
protoNum, err := r.af.protoNum(rule.Protocol)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("convert protocol to number: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Request forwarding before queueing rules: addDnatRedirect/addDnatMasq
|
|
||||||
// buffer netlink messages on r.conn that the next caller's Flush would
|
|
||||||
// commit if we returned without flushing them ourselves.
|
|
||||||
if err := r.ipFwdState.RequestForwarding(r.isV6()); err != nil {
|
|
||||||
return nil, fmt.Errorf("enable forwarding: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := r.addDnatRedirect(rule, protoNum, ruleID); err != nil {
|
|
||||||
r.releaseForwarding()
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := r.addDnatMasq(rule, protoNum, ruleID); err != nil {
|
|
||||||
r.releaseForwarding()
|
|
||||||
delete(r.rules, ruleID+dnatSuffix)
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Unlike iptables, there's no point in adding "out" rules in the forward chain here as our policy is ACCEPT.
|
|
||||||
// To overcome DROP policies in other chains, we'd have to add rules to the chains there.
|
|
||||||
// We also cannot just add "oif <iface> accept" there and filter in our own table as we don't know what is supposed to be allowed.
|
|
||||||
// TODO: find chains with drop policies and add rules there
|
|
||||||
|
|
||||||
if err := r.conn.Flush(); err != nil {
|
|
||||||
r.releaseForwarding()
|
|
||||||
delete(r.rules, ruleID+dnatSuffix)
|
|
||||||
delete(r.rules, ruleID+snatSuffix)
|
|
||||||
return nil, fmt.Errorf("flush rules: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return &rule, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *family) addDnatRedirect(rule firewall.ForwardRule, protoNum uint8, ruleID firewall.RuleID) error {
|
|
||||||
dnatExprs := []expr.Any{
|
|
||||||
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpNeq,
|
|
||||||
Register: 1,
|
|
||||||
Data: ifname(r.wgIface.Name()),
|
|
||||||
},
|
|
||||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Register: 1,
|
|
||||||
Data: []byte{protoNum},
|
|
||||||
},
|
|
||||||
&expr.Payload{
|
|
||||||
DestRegister: 1,
|
|
||||||
Base: expr.PayloadBaseTransportHeader,
|
|
||||||
Offset: 2,
|
|
||||||
Len: 2,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
portExprs, err := r.applyPort(&rule.DestinationPort, false)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("apply destination port: %w", err)
|
|
||||||
}
|
|
||||||
dnatExprs = append(dnatExprs, portExprs...)
|
|
||||||
|
|
||||||
// shifted translated port is not supported in nftables, so we hand this over to xtables
|
|
||||||
if rule.TranslatedPort.IsRange && len(rule.TranslatedPort.Values) == 2 {
|
|
||||||
if rule.TranslatedPort.Values[0] != rule.DestinationPort.Values[0] ||
|
|
||||||
rule.TranslatedPort.Values[1] != rule.DestinationPort.Values[1] {
|
|
||||||
return r.addXTablesRedirect(dnatExprs, ruleID, rule)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
additionalExprs, regProtoMin, regProtoMax, err := r.handleTranslatedPort(rule)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
dnatExprs = append(dnatExprs, additionalExprs...)
|
|
||||||
|
|
||||||
dnatExprs = append(dnatExprs,
|
|
||||||
&expr.NAT{
|
|
||||||
Type: expr.NATTypeDestNAT,
|
|
||||||
Family: uint32(r.af.tableFamily),
|
|
||||||
RegAddrMin: 1,
|
|
||||||
RegProtoMin: regProtoMin,
|
|
||||||
RegProtoMax: regProtoMax,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
dnatRule := &nftables.Rule{
|
|
||||||
Table: r.workTable,
|
|
||||||
Chain: r.chains[chainNameRoutingRdr],
|
|
||||||
Exprs: dnatExprs,
|
|
||||||
UserData: []byte(ruleID + dnatSuffix),
|
|
||||||
}
|
|
||||||
r.conn.AddRule(dnatRule)
|
|
||||||
r.rules[ruleID+dnatSuffix] = dnatRule
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *family) handleTranslatedPort(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
|
|
||||||
switch {
|
|
||||||
case rule.TranslatedPort.IsRange && len(rule.TranslatedPort.Values) == 2:
|
|
||||||
return r.handlePortRange(rule)
|
|
||||||
case len(rule.TranslatedPort.Values) == 0:
|
|
||||||
return r.handleAddressOnly(rule)
|
|
||||||
case len(rule.TranslatedPort.Values) == 1:
|
|
||||||
return r.handleSinglePort(rule)
|
|
||||||
default:
|
|
||||||
return nil, 0, 0, fmt.Errorf("invalid translated port: %v", rule.TranslatedPort)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *family) handlePortRange(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
|
|
||||||
exprs := []expr.Any{
|
|
||||||
&expr.Immediate{
|
|
||||||
Register: 1,
|
|
||||||
Data: rule.TranslatedAddress.AsSlice(),
|
|
||||||
},
|
|
||||||
&expr.Immediate{
|
|
||||||
Register: 2,
|
|
||||||
Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[0]),
|
|
||||||
},
|
|
||||||
&expr.Immediate{
|
|
||||||
Register: 3,
|
|
||||||
Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[1]),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
return exprs, 2, 3, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *family) handleAddressOnly(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
|
|
||||||
exprs := []expr.Any{
|
|
||||||
&expr.Immediate{
|
|
||||||
Register: 1,
|
|
||||||
Data: rule.TranslatedAddress.AsSlice(),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
return exprs, 0, 0, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *family) handleSinglePort(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
|
|
||||||
exprs := []expr.Any{
|
|
||||||
&expr.Immediate{
|
|
||||||
Register: 1,
|
|
||||||
Data: rule.TranslatedAddress.AsSlice(),
|
|
||||||
},
|
|
||||||
&expr.Immediate{
|
|
||||||
Register: 2,
|
|
||||||
Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[0]),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
return exprs, 2, 0, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *family) addXTablesRedirect(dnatExprs []expr.Any, ruleID firewall.RuleID, rule firewall.ForwardRule) error {
|
|
||||||
dnatExprs = append(dnatExprs,
|
|
||||||
&expr.Counter{},
|
|
||||||
&expr.Target{
|
|
||||||
Name: "DNAT",
|
|
||||||
Rev: 2,
|
|
||||||
Info: &xt.NatRange2{
|
|
||||||
NatRange: xt.NatRange{
|
|
||||||
Flags: uint(xt.NatRangeMapIPs | xt.NatRangeProtoSpecified | xt.NatRangeProtoOffset),
|
|
||||||
MinIP: rule.TranslatedAddress.AsSlice(),
|
|
||||||
MaxIP: rule.TranslatedAddress.AsSlice(),
|
|
||||||
MinPort: rule.TranslatedPort.Values[0],
|
|
||||||
MaxPort: rule.TranslatedPort.Values[1],
|
|
||||||
},
|
|
||||||
BasePort: rule.DestinationPort.Values[0],
|
|
||||||
},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
natTable := &nftables.Table{
|
|
||||||
Name: tableNat,
|
|
||||||
Family: r.af.tableFamily,
|
|
||||||
}
|
|
||||||
dnatRule := &nftables.Rule{
|
|
||||||
Table: natTable,
|
|
||||||
Chain: &nftables.Chain{
|
|
||||||
Name: chainNameNatPrerouting,
|
|
||||||
Table: natTable,
|
|
||||||
Type: nftables.ChainTypeNAT,
|
|
||||||
Hooknum: nftables.ChainHookPrerouting,
|
|
||||||
Priority: nftables.ChainPriorityNATDest,
|
|
||||||
},
|
|
||||||
Exprs: dnatExprs,
|
|
||||||
UserData: []byte(ruleID + dnatSuffix),
|
|
||||||
}
|
|
||||||
r.conn.AddRule(dnatRule)
|
|
||||||
r.rules[ruleID+dnatSuffix] = dnatRule
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *family) addDnatMasq(rule firewall.ForwardRule, protoNum uint8, ruleID firewall.RuleID) error {
|
|
||||||
portExprs, err := r.applyPort(&rule.TranslatedPort, false)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("apply translated port: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
masqExprs := []expr.Any{
|
|
||||||
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Register: 1,
|
|
||||||
Data: ifname(r.wgIface.Name()),
|
|
||||||
},
|
|
||||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Register: 1,
|
|
||||||
Data: []byte{protoNum},
|
|
||||||
},
|
|
||||||
&expr.Payload{
|
|
||||||
DestRegister: 1,
|
|
||||||
Base: expr.PayloadBaseNetworkHeader,
|
|
||||||
Offset: r.af.dstAddrOffset,
|
|
||||||
Len: r.af.addrLen,
|
|
||||||
},
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Register: 1,
|
|
||||||
Data: rule.TranslatedAddress.AsSlice(),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
masqExprs = append(masqExprs, portExprs...)
|
|
||||||
masqExprs = append(masqExprs, &expr.Masq{})
|
|
||||||
|
|
||||||
masqRule := &nftables.Rule{
|
|
||||||
Table: r.workTable,
|
|
||||||
Chain: r.chains[chainNameRoutingNat],
|
|
||||||
Exprs: masqExprs,
|
|
||||||
UserData: []byte(ruleID + snatSuffix),
|
|
||||||
}
|
|
||||||
r.conn.AddRule(masqRule)
|
|
||||||
r.rules[ruleID+snatSuffix] = masqRule
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *family) DeleteDNATRule(rule firewall.Rule) error {
|
|
||||||
ruleID := rule.ID()
|
|
||||||
|
|
||||||
if err := r.refreshRulesMap(); err != nil {
|
|
||||||
return fmt.Errorf(refreshRulesMapError, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var merr *multierror.Error
|
|
||||||
var needsFlush bool
|
|
||||||
var found bool
|
|
||||||
|
|
||||||
if dnatRule, exists := r.rules[ruleID+dnatSuffix]; exists {
|
|
||||||
found = true
|
|
||||||
if dnatRule.Handle == 0 {
|
|
||||||
log.Warnf("dnat rule %s has no handle, removing stale entry", ruleID+dnatSuffix)
|
|
||||||
delete(r.rules, ruleID+dnatSuffix)
|
|
||||||
} else if err := r.conn.DelRule(dnatRule); err != nil {
|
|
||||||
merr = multierror.Append(merr, fmt.Errorf("delete dnat rule: %w", err))
|
|
||||||
} else {
|
|
||||||
needsFlush = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if masqRule, exists := r.rules[ruleID+snatSuffix]; exists {
|
|
||||||
found = true
|
|
||||||
if masqRule.Handle == 0 {
|
|
||||||
log.Warnf("snat rule %s has no handle, removing stale entry", ruleID+snatSuffix)
|
|
||||||
delete(r.rules, ruleID+snatSuffix)
|
|
||||||
} else if err := r.conn.DelRule(masqRule); err != nil {
|
|
||||||
merr = multierror.Append(merr, fmt.Errorf("delete snat rule: %w", err))
|
|
||||||
} else {
|
|
||||||
needsFlush = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if needsFlush {
|
|
||||||
if err := r.conn.Flush(); err != nil {
|
|
||||||
merr = multierror.Append(merr, fmt.Errorf(flushError, err))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if merr != nil {
|
|
||||||
return nberrors.FormatErrorOrNil(merr)
|
|
||||||
}
|
|
||||||
|
|
||||||
delete(r.rules, ruleID+dnatSuffix)
|
|
||||||
delete(r.rules, ruleID+snatSuffix)
|
|
||||||
|
|
||||||
// Release once, only if the rule was present and removed.
|
|
||||||
if found {
|
|
||||||
r.releaseForwarding()
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// releaseForwarding drops one IP forwarding reference, logging any error.
|
|
||||||
func (r *family) releaseForwarding() {
|
|
||||||
if err := r.ipFwdState.ReleaseForwarding(r.isV6()); err != nil {
|
|
||||||
log.Errorf("release IP forwarding: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// isV6 reports whether this family handles the IPv6 table.
|
|
||||||
func (r *family) isV6() bool {
|
|
||||||
return r.af.tableFamily == nftables.TableFamilyIPv6
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *family) AddInboundDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
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))
|
ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
||||||
|
|
||||||
|
|||||||
@@ -1,249 +0,0 @@
|
|||||||
//go:build privileged
|
|
||||||
|
|
||||||
package nftables
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
fw "github.com/netbirdio/netbird/client/firewall/manager"
|
|
||||||
"github.com/netbirdio/netbird/client/iface"
|
|
||||||
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
|
||||||
)
|
|
||||||
|
|
||||||
func nftRefcountIfaceV4() *iFaceMock {
|
|
||||||
return &iFaceMock{
|
|
||||||
NameFunc: func() string { return "wt-refcount" },
|
|
||||||
AddressFunc: func() wgaddr.Address {
|
|
||||||
return wgaddr.Address{
|
|
||||||
IP: netip.MustParseAddr("100.96.0.1"),
|
|
||||||
Network: netip.MustParsePrefix("100.96.0.0/16"),
|
|
||||||
}
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func nftRefcountIfaceDual() *iFaceMock {
|
|
||||||
return &iFaceMock{
|
|
||||||
NameFunc: func() string { return "wt-refcount" },
|
|
||||||
AddressFunc: func() wgaddr.Address {
|
|
||||||
return wgaddr.Address{
|
|
||||||
IP: netip.MustParseAddr("100.96.0.1"),
|
|
||||||
Network: netip.MustParsePrefix("100.96.0.0/16"),
|
|
||||||
IPv6: netip.MustParseAddr("fd00::1"),
|
|
||||||
IPv6Net: netip.MustParsePrefix("fd00::/64"),
|
|
||||||
}
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newNftRefcountManager(t *testing.T, dual bool) *Manager {
|
|
||||||
t.Helper()
|
|
||||||
if check() != NFTABLES {
|
|
||||||
t.Skip("nftables not supported on this system")
|
|
||||||
}
|
|
||||||
var ifMock *iFaceMock
|
|
||||||
if dual {
|
|
||||||
ifMock = nftRefcountIfaceDual()
|
|
||||||
} else {
|
|
||||||
ifMock = nftRefcountIfaceV4()
|
|
||||||
}
|
|
||||||
m, err := Create(ifMock, iface.DefaultMTU)
|
|
||||||
require.NoError(t, err, "create manager")
|
|
||||||
require.NoError(t, m.Init(nil), "init manager")
|
|
||||||
t.Cleanup(func() {
|
|
||||||
require.NoError(t, m.Close(nil), "close manager")
|
|
||||||
})
|
|
||||||
return m
|
|
||||||
}
|
|
||||||
|
|
||||||
func dnatV4(port uint16) fw.ForwardRule {
|
|
||||||
return fw.ForwardRule{
|
|
||||||
Protocol: fw.ProtocolTCP,
|
|
||||||
DestinationPort: fw.Port{Values: []uint16{port}},
|
|
||||||
TranslatedAddress: netip.MustParseAddr("100.96.0.2"),
|
|
||||||
TranslatedPort: fw.Port{Values: []uint16{80}},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func dnatV6(port uint16) fw.ForwardRule {
|
|
||||||
return fw.ForwardRule{
|
|
||||||
Protocol: fw.ProtocolTCP,
|
|
||||||
DestinationPort: fw.Port{Values: []uint16{port}},
|
|
||||||
TranslatedAddress: netip.MustParseAddr("fd00::2"),
|
|
||||||
TranslatedPort: fw.Port{Values: []uint16{80}},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNftablesDNAT_RefcountBalancedV4 verifies that Add/Delete pairs leave the
|
|
||||||
// v4 refcount at zero.
|
|
||||||
func TestNftablesDNAT_RefcountBalancedV4(t *testing.T) {
|
|
||||||
m := newNftRefcountManager(t, false)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(dnatV4(8081))
|
|
||||||
require.NoError(t, err, "add v4 dnat 1")
|
|
||||||
v4, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 1, v4, "v4 refcount after first add")
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount unchanged")
|
|
||||||
|
|
||||||
r2, err := m.AddDNATRule(dnatV4(8082))
|
|
||||||
require.NoError(t, err, "add v4 dnat 2")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 2, v4, "v4 refcount after second add")
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount unchanged")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1), "delete v4 dnat 1")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 1, v4, "v4 refcount after first delete")
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount unchanged")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r2), "delete v4 dnat 2")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4, "v4 refcount after second delete")
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount unchanged")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNftablesDNAT_RefcountBalancedV6 verifies the v6 path increments v6 only
|
|
||||||
// and decrements back to zero on Delete.
|
|
||||||
func TestNftablesDNAT_RefcountBalancedV6(t *testing.T) {
|
|
||||||
m := newNftRefcountManager(t, true)
|
|
||||||
require.NotNil(t, m.family6, "v6 family")
|
|
||||||
require.Same(t, m.family4.ipFwdState, m.family6.ipFwdState, "shared state")
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(dnatV6(9091))
|
|
||||||
require.NoError(t, err, "add v6 dnat 1")
|
|
||||||
v4, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 0, v4, "v4 refcount unchanged")
|
|
||||||
assert.Equal(t, 1, v6, "v6 refcount after first add")
|
|
||||||
|
|
||||||
r2, err := m.AddDNATRule(dnatV6(9092))
|
|
||||||
require.NoError(t, err, "add v6 dnat 2")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4)
|
|
||||||
assert.Equal(t, 2, v6, "v6 refcount after second add")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat 1")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4, "v4 refcount unchanged")
|
|
||||||
assert.Equal(t, 1, v6, "v6 refcount after first delete")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r2), "delete v6 dnat 2")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4)
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount after second delete")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNftablesDNAT_DuplicateAddNoLeak verifies that a duplicate Add (same
|
|
||||||
// ForwardRule) does not double-increment the refcount.
|
|
||||||
func TestNftablesDNAT_DuplicateAddNoLeak(t *testing.T) {
|
|
||||||
m := newNftRefcountManager(t, true)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
rule := dnatV4(8083)
|
|
||||||
r1, err := m.AddDNATRule(rule)
|
|
||||||
require.NoError(t, err, "add v4 dnat")
|
|
||||||
v4, _ := state.Counts()
|
|
||||||
assert.Equal(t, 1, v4)
|
|
||||||
|
|
||||||
// duplicate add: same rule ID, must be a no-op for the refcount.
|
|
||||||
_, err = m.AddDNATRule(rule)
|
|
||||||
require.NoError(t, err, "duplicate add")
|
|
||||||
v4, _ = state.Counts()
|
|
||||||
assert.Equal(t, 1, v4, "duplicate add must not increment")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1), "delete v4 dnat")
|
|
||||||
v4, _ = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4, "single delete must drop to zero")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNftablesDNAT_DeleteMissingNoUnderflow verifies deleting a rule that was
|
|
||||||
// never added does not underflow the refcount.
|
|
||||||
func TestNftablesDNAT_DeleteMissingNoUnderflow(t *testing.T) {
|
|
||||||
m := newNftRefcountManager(t, true)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
// Construct a Rule reference for something never added. The router stores
|
|
||||||
// rules by ID(), and DeleteDNATRule looks them up in r.rules; a missing
|
|
||||||
// entry must be a no-op rather than calling Release.
|
|
||||||
phantom := dnatV4(8099)
|
|
||||||
require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4 dnat")
|
|
||||||
v4, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 0, v4, "v4 refcount unaffected by missing delete")
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount unaffected")
|
|
||||||
|
|
||||||
phantom6 := dnatV6(9099)
|
|
||||||
require.NoError(t, m.DeleteDNATRule(&phantom6), "delete missing v6 dnat")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4)
|
|
||||||
assert.Equal(t, 0, v6, "v6 refcount unaffected by missing delete")
|
|
||||||
|
|
||||||
// And after a phantom delete, a real add still results in count=1.
|
|
||||||
r1, err := m.AddDNATRule(dnatV4(8100))
|
|
||||||
require.NoError(t, err, "add v4 dnat after phantom delete")
|
|
||||||
v4, _ = state.Counts()
|
|
||||||
assert.Equal(t, 1, v4, "real add still increments after phantom delete")
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1))
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNftablesRouting_RepeatedEnableSingleReference verifies that EnableRouting
|
|
||||||
// (called on every network-map update) holds at most one reference per family
|
|
||||||
// and a single DisableRouting drops both back to zero.
|
|
||||||
func TestNftablesRouting_RepeatedEnableSingleReference(t *testing.T) {
|
|
||||||
m := newNftRefcountManager(t, true)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
require.NoError(t, m.EnableRouting(), "first enable")
|
|
||||||
require.NoError(t, m.EnableRouting(), "second enable")
|
|
||||||
require.NoError(t, m.EnableRouting(), "third enable")
|
|
||||||
v4, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 1, v4, "repeated enable holds a single v4 reference")
|
|
||||||
assert.Equal(t, 1, v6, "repeated enable holds a single v6 reference")
|
|
||||||
|
|
||||||
require.NoError(t, m.DisableRouting(), "disable")
|
|
||||||
v4, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v4, "single disable releases the v4 reference")
|
|
||||||
assert.Equal(t, 0, v6, "single disable releases the v6 reference")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNftablesRouting_DisableKeepsDNATReference verifies that an unpaired
|
|
||||||
// DisableRouting does not release references held by active DNAT rules.
|
|
||||||
func TestNftablesRouting_DisableKeepsDNATReference(t *testing.T) {
|
|
||||||
m := newNftRefcountManager(t, true)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(dnatV6(9095))
|
|
||||||
require.NoError(t, err, "add v6 dnat")
|
|
||||||
|
|
||||||
require.NoError(t, m.DisableRouting(), "unpaired disable")
|
|
||||||
_, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting")
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat")
|
|
||||||
_, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v6, "delete releases the DNAT reference")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNftablesDNAT_DoubleDeleteNoUnderflow verifies that deleting the same rule
|
|
||||||
// twice does not underflow the refcount (the second delete is a no-op).
|
|
||||||
func TestNftablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) {
|
|
||||||
m := newNftRefcountManager(t, true)
|
|
||||||
state := m.family4.ipFwdState
|
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(dnatV6(9093))
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, v6 := state.Counts()
|
|
||||||
assert.Equal(t, 1, v6)
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1), "first delete")
|
|
||||||
_, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v6)
|
|
||||||
|
|
||||||
require.NoError(t, m.DeleteDNATRule(r1), "second delete must be no-op")
|
|
||||||
_, v6 = state.Counts()
|
|
||||||
assert.Equal(t, 0, v6, "double delete must not underflow")
|
|
||||||
}
|
|
||||||
@@ -24,7 +24,6 @@ const (
|
|||||||
tableRaw = "raw"
|
tableRaw = "raw"
|
||||||
tableSecurity = "security"
|
tableSecurity = "security"
|
||||||
|
|
||||||
chainNameNatPrerouting = "PREROUTING"
|
|
||||||
chainNameRoutingFw = "netbird-rt-fwd"
|
chainNameRoutingFw = "netbird-rt-fwd"
|
||||||
chainNameRoutingNat = "netbird-rt-postrouting"
|
chainNameRoutingNat = "netbird-rt-postrouting"
|
||||||
chainNameRoutingRdr = "netbird-rt-redirect"
|
chainNameRoutingRdr = "netbird-rt-redirect"
|
||||||
@@ -47,9 +46,6 @@ const (
|
|||||||
userDataAcceptForwardRuleOif = "frwacceptoif"
|
userDataAcceptForwardRuleOif = "frwacceptoif"
|
||||||
userDataAcceptInputRule = "inputaccept"
|
userDataAcceptInputRule = "inputaccept"
|
||||||
|
|
||||||
dnatSuffix firewall.RuleID = "_dnat"
|
|
||||||
snatSuffix firewall.RuleID = "_snat"
|
|
||||||
|
|
||||||
// ipv4TCPHeaderSize is the minimum IPv4 (20) + TCP (20) header size for MSS calculation.
|
// ipv4TCPHeaderSize is the minimum IPv4 (20) + TCP (20) header size for MSS calculation.
|
||||||
ipv4TCPHeaderSize = 40
|
ipv4TCPHeaderSize = 40
|
||||||
// ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation.
|
// 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)
|
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)
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -197,11 +197,6 @@ func (r *family) hasRule(id firewall.RuleID) bool {
|
|||||||
return ok
|
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
|
// DeleteFilterRule removes a previously installed filter rule. Source
|
||||||
// set references are recovered from the stored rule's expressions via
|
// set references are recovered from the stored rule's expressions via
|
||||||
// findSets and dropped from the shared refcounter.
|
// findSets and dropped from the shared refcounter.
|
||||||
|
|||||||
@@ -12,7 +12,6 @@ import (
|
|||||||
"github.com/google/nftables/expr"
|
"github.com/google/nftables/expr"
|
||||||
"github.com/hashicorp/go-multierror"
|
"github.com/hashicorp/go-multierror"
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
|
|
||||||
nberrors "github.com/netbirdio/netbird/client/errors"
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
@@ -55,9 +54,6 @@ type Manager struct {
|
|||||||
// IPv6 counterpart, nil when no v6 overlay.
|
// IPv6 counterpart, nil when no v6 overlay.
|
||||||
family6 *family
|
family6 *family
|
||||||
|
|
||||||
notrackOutputChain *nftables.Chain
|
|
||||||
notrackPreroutingChain *nftables.Chain
|
|
||||||
|
|
||||||
extMonitor *externalChainMonitor
|
extMonitor *externalChainMonitor
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -170,10 +166,6 @@ func (m *Manager) initFirewall() (err error) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := m.initNoTrackChains(workTable); err != nil {
|
|
||||||
log.Warnf("raw priority chains not available, notrack rules will be disabled: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -260,7 +252,7 @@ func (m *Manager) DeleteFilterRule(rule firewall.Rule) error {
|
|||||||
m.mutex.Lock()
|
m.mutex.Lock()
|
||||||
defer m.mutex.Unlock()
|
defer m.mutex.Unlock()
|
||||||
|
|
||||||
fam, err := m.familyForRuleID(rule.ID(), (*family).hasRule, false)
|
fam, err := m.familyForRuleID(rule.ID(), (*family).hasRule)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
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
|
// 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 supplied lookup, and falls back to the v4 family on a miss.
|
||||||
// the NAT/DNAT rule maps from the kernel once and re-checks before falling
|
func (m *Manager) familyForRuleID(id firewall.RuleID, has func(*family, firewall.RuleID) bool) (*family, error) {
|
||||||
// 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) {
|
|
||||||
if has(m.family4, id) {
|
if has(m.family4, id) {
|
||||||
return m.family4, nil
|
return m.family4, nil
|
||||||
}
|
}
|
||||||
@@ -282,18 +271,6 @@ func (m *Manager) familyForRuleID(id firewall.RuleID, has func(*family, firewall
|
|||||||
if has(m.family6, id) {
|
if has(m.family6, id) {
|
||||||
return m.family6, nil
|
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
|
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
|
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
|
// UpdateSet updates the set with the given prefixes
|
||||||
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
|
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
|
||||||
m.mutex.Lock()
|
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)
|
return m.family4.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
|
||||||
chainNameRawOutput = "netbird-raw-out"
|
|
||||||
chainNameRawPrerouting = "netbird-raw-pre"
|
|
||||||
)
|
|
||||||
|
|
||||||
// SetupEBPFProxyNoTrack creates notrack rules for eBPF proxy loopback traffic.
|
|
||||||
// This prevents conntrack from tracking WireGuard proxy traffic on loopback, which
|
|
||||||
// can interfere with MASQUERADE rules (e.g., from container runtimes like Podman/netavark).
|
|
||||||
//
|
|
||||||
// Traffic flows that need NOTRACK:
|
|
||||||
//
|
|
||||||
// 1. Egress: WireGuard -> fake endpoint (before eBPF rewrite)
|
|
||||||
// src=127.0.0.1:wgPort -> dst=127.0.0.1:fakePort
|
|
||||||
// Matched by: sport=wgPort
|
|
||||||
//
|
|
||||||
// 2. Egress: Proxy -> WireGuard (via raw socket)
|
|
||||||
// src=127.0.0.1:fakePort -> dst=127.0.0.1:wgPort
|
|
||||||
// Matched by: dport=wgPort
|
|
||||||
//
|
|
||||||
// 3. Ingress: Packets to WireGuard
|
|
||||||
// dst=127.0.0.1:wgPort
|
|
||||||
// Matched by: dport=wgPort
|
|
||||||
//
|
|
||||||
// 4. Ingress: Packets to proxy (after eBPF rewrite)
|
|
||||||
// dst=127.0.0.1:proxyPort
|
|
||||||
// Matched by: dport=proxyPort
|
|
||||||
//
|
|
||||||
// Rules are cleaned up when the firewall manager is closed.
|
|
||||||
func (m *Manager) SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error {
|
|
||||||
m.mutex.Lock()
|
|
||||||
defer m.mutex.Unlock()
|
|
||||||
|
|
||||||
if m.notrackOutputChain == nil || m.notrackPreroutingChain == nil {
|
|
||||||
return fmt.Errorf("notrack chains not initialized")
|
|
||||||
}
|
|
||||||
|
|
||||||
proxyPortBytes := binaryutil.BigEndian.PutUint16(proxyPort)
|
|
||||||
wgPortBytes := binaryutil.BigEndian.PutUint16(wgPort)
|
|
||||||
loopback := []byte{127, 0, 0, 1}
|
|
||||||
|
|
||||||
// Egress rules: match outgoing loopback UDP packets
|
|
||||||
m.rConn.AddRule(&nftables.Rule{
|
|
||||||
Table: m.notrackOutputChain.Table,
|
|
||||||
Chain: m.notrackOutputChain,
|
|
||||||
Exprs: []expr.Any{
|
|
||||||
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
|
||||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 2},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // sport=wgPort
|
|
||||||
&expr.Counter{},
|
|
||||||
&expr.Notrack{},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.rConn.AddRule(&nftables.Rule{
|
|
||||||
Table: m.notrackOutputChain.Table,
|
|
||||||
Chain: m.notrackOutputChain,
|
|
||||||
Exprs: []expr.Any{
|
|
||||||
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
|
||||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // dport=wgPort
|
|
||||||
&expr.Counter{},
|
|
||||||
&expr.Notrack{},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
// Ingress rules: match incoming loopback UDP packets
|
|
||||||
m.rConn.AddRule(&nftables.Rule{
|
|
||||||
Table: m.notrackPreroutingChain.Table,
|
|
||||||
Chain: m.notrackPreroutingChain,
|
|
||||||
Exprs: []expr.Any{
|
|
||||||
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
|
||||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // dport=wgPort
|
|
||||||
&expr.Counter{},
|
|
||||||
&expr.Notrack{},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.rConn.AddRule(&nftables.Rule{
|
|
||||||
Table: m.notrackPreroutingChain.Table,
|
|
||||||
Chain: m.notrackPreroutingChain,
|
|
||||||
Exprs: []expr.Any{
|
|
||||||
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
|
||||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
|
|
||||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
|
|
||||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: proxyPortBytes}, // dport=proxyPort
|
|
||||||
&expr.Counter{},
|
|
||||||
&expr.Notrack{},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
if err := m.rConn.Flush(); err != nil {
|
|
||||||
return fmt.Errorf("flush notrack rules: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Debugf("set up ebpf proxy notrack rules for ports %d,%d", proxyPort, wgPort)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Manager) initNoTrackChains(table *nftables.Table) error {
|
|
||||||
m.notrackOutputChain = m.rConn.AddChain(&nftables.Chain{
|
|
||||||
Name: chainNameRawOutput,
|
|
||||||
Table: table,
|
|
||||||
Type: nftables.ChainTypeFilter,
|
|
||||||
Hooknum: nftables.ChainHookOutput,
|
|
||||||
Priority: nftables.ChainPriorityRaw,
|
|
||||||
})
|
|
||||||
|
|
||||||
m.notrackPreroutingChain = m.rConn.AddChain(&nftables.Chain{
|
|
||||||
Name: chainNameRawPrerouting,
|
|
||||||
Table: table,
|
|
||||||
Type: nftables.ChainTypeFilter,
|
|
||||||
Hooknum: nftables.ChainHookPrerouting,
|
|
||||||
Priority: nftables.ChainPriorityRaw,
|
|
||||||
})
|
|
||||||
|
|
||||||
if err := m.rConn.Flush(); err != nil {
|
|
||||||
return fmt.Errorf("flush chain creation: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Manager) refreshNoTrackChains() error {
|
|
||||||
chains, err := m.rConn.ListChainsOfTableFamily(nftables.TableFamilyIPv4)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("list chains: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
tableName := getTableName()
|
|
||||||
for _, c := range chains {
|
|
||||||
if c.Table.Name != tableName {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
switch c.Name {
|
|
||||||
case chainNameRawOutput:
|
|
||||||
m.notrackOutputChain = c
|
|
||||||
case chainNameRawPrerouting:
|
|
||||||
m.notrackPreroutingChain = c
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Manager) createWorkTable() (*nftables.Table, error) {
|
func (m *Manager) createWorkTable() (*nftables.Table, error) {
|
||||||
return m.createWorkTableFamily(nftables.TableFamilyIPv4)
|
return m.createWorkTableFamily(nftables.TableFamilyIPv4)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -378,18 +378,6 @@ func TestNftablesManagerCompatibilityWithIptables(t *testing.T) {
|
|||||||
err = manager.AddNatRule(pair)
|
err = manager.AddNatRule(pair)
|
||||||
require.NoError(t, err, "failed to add NAT rule")
|
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)
|
stdout, stderr = runIptablesSave(t)
|
||||||
verifyIptablesOutput(t, stdout, stderr)
|
verifyIptablesOutput(t, stdout, stderr)
|
||||||
}
|
}
|
||||||
@@ -453,18 +441,6 @@ func TestNftablesManagerIPv6CompatibilityWithIp6tables(t *testing.T) {
|
|||||||
})
|
})
|
||||||
require.NoError(t, err, "add v6 NAT rule")
|
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)
|
stdout, stderr := runIptablesSave(t)
|
||||||
verifyIptablesOutput(t, stdout, stderr)
|
verifyIptablesOutput(t, stdout, stderr)
|
||||||
|
|
||||||
|
|||||||
@@ -192,7 +192,7 @@ func (r *family) addPostroutingRules() {
|
|||||||
Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkMasquerade),
|
Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkMasquerade),
|
||||||
},
|
},
|
||||||
|
|
||||||
// We need to exclude the loopback interface as this changes the ebpf proxy port
|
// We need to exclude the loopback interface as this changes the wg proxy port
|
||||||
&expr.Meta{
|
&expr.Meta{
|
||||||
Key: expr.MetaKeyOIFNAME,
|
Key: expr.MetaKeyOIFNAME,
|
||||||
Register: 1,
|
Register: 1,
|
||||||
@@ -459,41 +459,6 @@ func (r *family) RemoveAllLegacyRouteRules() error {
|
|||||||
return nberrors.FormatErrorOrNil(merr)
|
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 {
|
func (r *family) RemoveNatRule(pair firewall.RouterPair) error {
|
||||||
if err := r.refreshRulesMap(); err != nil {
|
if err := r.refreshRulesMap(); err != nil {
|
||||||
return fmt.Errorf(refreshRulesMapError, err)
|
return fmt.Errorf(refreshRulesMapError, err)
|
||||||
|
|||||||
@@ -879,12 +879,6 @@ func (m *Manager) resetState() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetupEBPFProxyNoTrack is not supported by the userspace firewall: eBPF isn't
|
|
||||||
// used in userspace mode, so this should never be called.
|
|
||||||
func (m *Manager) SetupEBPFProxyNoTrack(uint16, uint16) error {
|
|
||||||
return errNotSupported
|
|
||||||
}
|
|
||||||
|
|
||||||
// UpdateSet updates the rule destinations associated with the given set
|
// UpdateSet updates the rule destinations associated with the given set
|
||||||
// by merging the existing prefixes with the new ones, then deduplicating.
|
// by merging the existing prefixes with the new ones, then deduplicating.
|
||||||
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
|
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
nberrors "github.com/netbirdio/netbird/client/errors"
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/wincmd"
|
||||||
)
|
)
|
||||||
|
|
||||||
type action string
|
type action string
|
||||||
@@ -91,7 +92,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err
|
|||||||
if action == addRule {
|
if action == addRule {
|
||||||
args = append(args, extraArgs...)
|
args = append(args, extraArgs...)
|
||||||
}
|
}
|
||||||
netshCmd := GetSystem32Command("netsh")
|
netshCmd := wincmd.System32("netsh")
|
||||||
cmd := exec.Command(netshCmd, args...)
|
cmd := exec.Command(netshCmd, args...)
|
||||||
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
|
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
|
||||||
return cmd.Run()
|
return cmd.Run()
|
||||||
@@ -100,7 +101,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err
|
|||||||
func isWindowsFirewallReachable() bool {
|
func isWindowsFirewallReachable() bool {
|
||||||
args := []string{"advfirewall", "show", "allprofiles", "state"}
|
args := []string{"advfirewall", "show", "allprofiles", "state"}
|
||||||
|
|
||||||
netshCmd := GetSystem32Command("netsh")
|
netshCmd := wincmd.System32("netsh")
|
||||||
|
|
||||||
cmd := exec.Command(netshCmd, args...)
|
cmd := exec.Command(netshCmd, args...)
|
||||||
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
|
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
|
||||||
@@ -117,23 +118,10 @@ func isWindowsFirewallReachable() bool {
|
|||||||
func isFirewallRuleActive(ruleName string) bool {
|
func isFirewallRuleActive(ruleName string) bool {
|
||||||
args := []string{"advfirewall", "firewall", "show", "rule", "name=" + ruleName}
|
args := []string{"advfirewall", "firewall", "show", "rule", "name=" + ruleName}
|
||||||
|
|
||||||
netshCmd := GetSystem32Command("netsh")
|
netshCmd := wincmd.System32("netsh")
|
||||||
|
|
||||||
cmd := exec.Command(netshCmd, args...)
|
cmd := exec.Command(netshCmd, args...)
|
||||||
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
|
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
|
||||||
_, err := cmd.Output()
|
_, err := cmd.Output()
|
||||||
return err == nil
|
return err == nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetSystem32Command checks if a command can be found in the system path and returns it. In case it can't find it
|
|
||||||
// in the path it will return the full path of a command assuming C:\windows\system32 as the base path.
|
|
||||||
func GetSystem32Command(command string) string {
|
|
||||||
_, err := exec.LookPath(command)
|
|
||||||
if err == nil {
|
|
||||||
return command
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Tracef("Command %s not found in PATH, using C:\\windows\\system32\\%s.exe path", command, command)
|
|
||||||
|
|
||||||
return "C:\\windows\\system32\\" + command + ".exe"
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -486,16 +486,6 @@ func incrementalUpdate(oldChecksum uint16, oldBytes, newBytes []byte) uint16 {
|
|||||||
return ^uint16(sum)
|
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.
|
// addPortRedirection adds a port redirection rule.
|
||||||
func (m *Manager) addPortRedirection(targetIP netip.Addr, protocol gopacket.LayerType, originalPort, translatedPort uint16) error {
|
func (m *Manager) addPortRedirection(targetIP netip.Addr, protocol gopacket.LayerType, originalPort, translatedPort uint16) error {
|
||||||
m.portDNATMutex.Lock()
|
m.portDNATMutex.Lock()
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -19,12 +19,18 @@ func buildPresharedKeyConfig(peerKey wgtypes.Key, psk wgtypes.Key, updateOnly bo
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// prefixesToIPNets converts prefixes on their way to a device. It is the only place that
|
||||||
|
// conversion happens, so it also normalizes: the device is then given the same form the
|
||||||
|
// store records, and a v4-mapped prefix cannot reach net.IPNet, which prints such an
|
||||||
|
// address as v4 while taking the length from its 16 byte mask and so turns
|
||||||
|
// ::ffff:10.1.2.3/64 into 10.1.2.3/0 — an allowed IP matching every v4 address.
|
||||||
func prefixesToIPNets(prefixes []netip.Prefix) []net.IPNet {
|
func prefixesToIPNets(prefixes []netip.Prefix) []net.IPNet {
|
||||||
ipNets := make([]net.IPNet, len(prefixes))
|
ipNets := make([]net.IPNet, len(prefixes))
|
||||||
for i, prefix := range prefixes {
|
for i, prefix := range prefixes {
|
||||||
|
normalized := normalizePrefix(prefix)
|
||||||
ipNets[i] = net.IPNet{
|
ipNets[i] = net.IPNet{
|
||||||
IP: prefix.Addr().AsSlice(), // Convert netip.Addr to net.IP
|
IP: normalized.Addr().AsSlice(),
|
||||||
Mask: net.CIDRMask(prefix.Bits(), prefix.Addr().BitLen()), // Create subnet mask
|
Mask: net.CIDRMask(normalized.Bits(), normalized.Addr().BitLen()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return ipNets
|
return ipNets
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"slices"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
@@ -18,16 +19,22 @@ import (
|
|||||||
type KernelConfigurer struct {
|
type KernelConfigurer struct {
|
||||||
deviceName string
|
deviceName string
|
||||||
statsCache *statsCache
|
statsCache *statsCache
|
||||||
|
allowedIPs *allowedIPStore
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// NewKernelConfigurer creates a configurer with an empty allowed IP mirror
|
||||||
|
// and a statistics cache for the named kernel device.
|
||||||
func NewKernelConfigurer(deviceName string) *KernelConfigurer {
|
func NewKernelConfigurer(deviceName string) *KernelConfigurer {
|
||||||
c := &KernelConfigurer{
|
c := &KernelConfigurer{
|
||||||
deviceName: deviceName,
|
deviceName: deviceName,
|
||||||
|
allowedIPs: newAllowedIPStore(),
|
||||||
}
|
}
|
||||||
c.statsCache = newStatsCache(statsCacheTTL, c.fetchStats)
|
c.statsCache = newStatsCache(statsCacheTTL, c.fetchStats)
|
||||||
return c
|
return c
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ConfigureInterface sets the device key, port and firewall mark, replacing all peers.
|
||||||
|
// The allowed IP mirror is reset only after the device accepts the configuration.
|
||||||
func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error {
|
func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error {
|
||||||
log.Debugf("adding Wireguard private key")
|
log.Debugf("adding Wireguard private key")
|
||||||
key, err := wgtypes.ParseKey(privateKey)
|
key, err := wgtypes.ParseKey(privateKey)
|
||||||
@@ -46,6 +53,8 @@ func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf(`received error "%w" while configuring interface %s with port %d`, err, c.deviceName, port)
|
return fmt.Errorf(`received error "%w" while configuring interface %s with port %d`, err, c.deviceName, port)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
c.allowedIPs.reset()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -58,9 +67,20 @@ func (c *KernelConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, upda
|
|||||||
}
|
}
|
||||||
|
|
||||||
cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly)
|
cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly)
|
||||||
return c.configure(cfg)
|
if err := c.configure(cfg); err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Without updateOnly this creates the peer when it is absent, so the store has to
|
||||||
|
// know about it even though no allowed IP was configured.
|
||||||
|
if !updateOnly {
|
||||||
|
c.allowedIPs.ensure(parsedPeerKey)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set.
|
||||||
|
// Prefixes assigned to this peer are transferred from their previous owners.
|
||||||
func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error {
|
func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error {
|
||||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -83,19 +103,23 @@ func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix,
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf(`received error "%w" while updating peer on interface %s with settings: allowed ips %s, endpoint %s`, err, c.deviceName, allowedIps, endpoint.String())
|
return fmt.Errorf(`received error "%w" while updating peer on interface %s with settings: allowed ips %s, endpoint %s`, err, c.deviceName, allowedIps, endpoint.String())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
c.allowedIPs.add(peerKeyParsed, allowedIps)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured.
|
||||||
|
// Neither the netlink API nor the userspace one can clear an endpoint in place, so the peer
|
||||||
|
// is removed and re-added with the allowed IPs it already had.
|
||||||
func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error {
|
func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error {
|
||||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get the existing peer to preserve its allowed IPs
|
allowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
|
||||||
existingPeer, err := c.getPeer(c.deviceName, peerKey)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get peer: %w", err)
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
removePeerCfg := wgtypes.PeerConfig{
|
removePeerCfg := wgtypes.PeerConfig{
|
||||||
@@ -104,26 +128,27 @@ func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{removePeerCfg}}); err != nil {
|
if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{removePeerCfg}}); err != nil {
|
||||||
return fmt.Errorf(`error removing peer %s from interface %s: %w`, peerKey, c.deviceName, err)
|
return fmt.Errorf("remove peer %s from interface %s: %w", peerKey, c.deviceName, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
//Re-add the peer without the endpoint but same AllowedIPs
|
|
||||||
reAddPeerCfg := wgtypes.PeerConfig{
|
reAddPeerCfg := wgtypes.PeerConfig{
|
||||||
PublicKey: peerKeyParsed,
|
PublicKey: peerKeyParsed,
|
||||||
AllowedIPs: existingPeer.AllowedIPs,
|
AllowedIPs: prefixesToIPNets(allowedIPs),
|
||||||
ReplaceAllowedIPs: true,
|
ReplaceAllowedIPs: true,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{reAddPeerCfg}}); err != nil {
|
if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{reAddPeerCfg}}); err != nil {
|
||||||
|
c.allowedIPs.forget(peerKeyParsed)
|
||||||
return fmt.Errorf(
|
return fmt.Errorf(
|
||||||
`error re-adding peer %s to interface %s with allowed IPs %v: %w`,
|
"re-add peer %s to interface %s with allowed IPs %v: %w",
|
||||||
peerKey, c.deviceName, existingPeer.AllowedIPs, err,
|
peerKey, c.deviceName, allowedIPs, err,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RemovePeer removes a peer and forgets its allowed IPs after a successful device write.
|
||||||
func (c *KernelConfigurer) RemovePeer(peerKey string) error {
|
func (c *KernelConfigurer) RemovePeer(peerKey string) error {
|
||||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -142,15 +167,13 @@ func (c *KernelConfigurer) RemovePeer(peerKey string) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf(`received error "%w" while removing peer %s from interface %s`, err, peerKey, c.deviceName)
|
return fmt.Errorf(`received error "%w" while removing peer %s from interface %s`, err, peerKey, c.deviceName)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
c.allowedIPs.forget(peerKeyParsed)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op.
|
||||||
func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
||||||
ipNet := net.IPNet{
|
|
||||||
IP: allowedIP.Addr().AsSlice(),
|
|
||||||
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
|
|
||||||
}
|
|
||||||
|
|
||||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -159,7 +182,7 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix)
|
|||||||
PublicKey: peerKeyParsed,
|
PublicKey: peerKeyParsed,
|
||||||
UpdateOnly: true,
|
UpdateOnly: true,
|
||||||
ReplaceAllowedIPs: false,
|
ReplaceAllowedIPs: false,
|
||||||
AllowedIPs: []net.IPNet{ipNet},
|
AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}),
|
||||||
}
|
}
|
||||||
|
|
||||||
config := wgtypes.Config{
|
config := wgtypes.Config{
|
||||||
@@ -169,52 +192,69 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix)
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf(`received error "%w" while adding allowed Ip to peer on interface %s with settings: allowed ips %s`, err, c.deviceName, allowedIP)
|
return fmt.Errorf(`received error "%w" while adding allowed Ip to peer on interface %s with settings: allowed ips %s`, err, c.deviceName, allowedIP)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP})
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RemoveAllowedIP removes a prefix while preserving the peer's other allowed IPs.
|
||||||
|
// A prefix not assigned to the peer is a no-op.
|
||||||
func (c *KernelConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
func (c *KernelConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
||||||
ipNet := net.IPNet{
|
|
||||||
IP: allowedIP.Addr().AsSlice(),
|
|
||||||
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
|
|
||||||
}
|
|
||||||
|
|
||||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("parse peer key: %w", err)
|
return fmt.Errorf("parse peer key: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
existingPeer, err := c.getPeer(c.deviceName, peerKey)
|
currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get peer: %w", err)
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
newAllowedIPs := existingPeer.AllowedIPs
|
idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP))
|
||||||
|
if idx < 0 {
|
||||||
for i, existingAllowedIP := range existingPeer.AllowedIPs {
|
return nil
|
||||||
if existingAllowedIP.String() == ipNet.String() {
|
|
||||||
newAllowedIPs = append(existingPeer.AllowedIPs[:i], existingPeer.AllowedIPs[i+1:]...) //nolint:gocritic
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1)
|
||||||
|
|
||||||
peer := wgtypes.PeerConfig{
|
peer := wgtypes.PeerConfig{
|
||||||
PublicKey: peerKeyParsed,
|
PublicKey: peerKeyParsed,
|
||||||
UpdateOnly: true,
|
UpdateOnly: true,
|
||||||
ReplaceAllowedIPs: true,
|
ReplaceAllowedIPs: true,
|
||||||
AllowedIPs: newAllowedIPs,
|
AllowedIPs: prefixesToIPNets(newAllowedIPs),
|
||||||
}
|
}
|
||||||
|
|
||||||
config := wgtypes.Config{
|
config := wgtypes.Config{
|
||||||
Peers: []wgtypes.PeerConfig{peer},
|
Peers: []wgtypes.PeerConfig{peer},
|
||||||
}
|
}
|
||||||
err = c.configure(config)
|
if err := c.configure(config); err != nil {
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("remove allowed IP %s on interface %s: %w", allowedIP, c.deviceName, err)
|
return fmt.Errorf("remove allowed IP %s on interface %s: %w", allowedIP, c.deviceName, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
c.allowedIPs.set(peerKeyParsed, newAllowedIPs)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) {
|
// peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device
|
||||||
|
// only for a peer the store has not seen. Dumping the device costs a netlink round trip
|
||||||
|
// proportional to the whole network map, and this runs on every relay and ICE transition.
|
||||||
|
func (c *KernelConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) {
|
||||||
|
if prefixes, ok := c.allowedIPs.get(peerKey); ok {
|
||||||
|
return prefixes, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
existingPeer, err := c.getPeer(c.deviceName, peerKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("get peer: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
prefixes := ipNetsToPrefixes(existingPeer.AllowedIPs)
|
||||||
|
c.allowedIPs.set(peerKey, prefixes)
|
||||||
|
return prefixes, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// getPeer scans the device for one peer. wgtypes.Key is an array, so the comparison is a
|
||||||
|
// plain equality: Key.String would base64 encode into a fresh allocation for every peer.
|
||||||
|
func (c *KernelConfigurer) getPeer(ifaceName string, peerPubKey wgtypes.Key) (wgtypes.Peer, error) {
|
||||||
wg, err := wgctrl.New()
|
wg, err := wgctrl.New()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return wgtypes.Peer{}, fmt.Errorf("wgctl: %w", err)
|
return wgtypes.Peer{}, fmt.Errorf("wgctl: %w", err)
|
||||||
@@ -231,7 +271,7 @@ func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer,
|
|||||||
return wgtypes.Peer{}, fmt.Errorf("get device %s: %w", ifaceName, err)
|
return wgtypes.Peer{}, fmt.Errorf("get device %s: %w", ifaceName, err)
|
||||||
}
|
}
|
||||||
for _, peer := range wgDevice.Peers {
|
for _, peer := range wgDevice.Peers {
|
||||||
if peer.PublicKey.String() == peerPubKey {
|
if peer.PublicKey == peerPubKey {
|
||||||
return peer, nil
|
return peer, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+117
-89
@@ -8,6 +8,7 @@ import (
|
|||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"runtime"
|
"runtime"
|
||||||
|
"slices"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -41,31 +42,38 @@ type WGUSPConfigurer struct {
|
|||||||
deviceName string
|
deviceName string
|
||||||
activityRecorder *bind.ActivityRecorder
|
activityRecorder *bind.ActivityRecorder
|
||||||
statsCache *statsCache
|
statsCache *statsCache
|
||||||
|
allowedIPs *allowedIPStore
|
||||||
|
|
||||||
uapiListener net.Listener
|
uapiListener net.Listener
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// NewUSPConfigurer creates a userspace configurer and starts its UAPI listener.
|
||||||
func NewUSPConfigurer(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
|
func NewUSPConfigurer(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
|
||||||
wgCfg := &WGUSPConfigurer{
|
wgCfg := &WGUSPConfigurer{
|
||||||
device: device,
|
device: device,
|
||||||
deviceName: deviceName,
|
deviceName: deviceName,
|
||||||
activityRecorder: activityRecorder,
|
activityRecorder: activityRecorder,
|
||||||
|
allowedIPs: newAllowedIPStore(),
|
||||||
}
|
}
|
||||||
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
|
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
|
||||||
wgCfg.startUAPI()
|
wgCfg.startUAPI()
|
||||||
return wgCfg
|
return wgCfg
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// NewUSPConfigurerNoUAPI creates a userspace configurer without a UAPI listener.
|
||||||
func NewUSPConfigurerNoUAPI(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
|
func NewUSPConfigurerNoUAPI(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
|
||||||
wgCfg := &WGUSPConfigurer{
|
wgCfg := &WGUSPConfigurer{
|
||||||
device: device,
|
device: device,
|
||||||
deviceName: deviceName,
|
deviceName: deviceName,
|
||||||
activityRecorder: activityRecorder,
|
activityRecorder: activityRecorder,
|
||||||
|
allowedIPs: newAllowedIPStore(),
|
||||||
}
|
}
|
||||||
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
|
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
|
||||||
return wgCfg
|
return wgCfg
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ConfigureInterface sets the device key, port and firewall mark, replacing all peers.
|
||||||
|
// The allowed IP mirror is reset only after the device accepts the configuration.
|
||||||
func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error {
|
func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error {
|
||||||
log.Debugf("adding Wireguard private key")
|
log.Debugf("adding Wireguard private key")
|
||||||
key, err := wgtypes.ParseKey(privateKey)
|
key, err := wgtypes.ParseKey(privateKey)
|
||||||
@@ -80,7 +88,12 @@ func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error
|
|||||||
ListenPort: &port,
|
ListenPort: &port,
|
||||||
}
|
}
|
||||||
|
|
||||||
return c.device.IpcSet(toWgUserspaceString(config))
|
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
c.allowedIPs.reset()
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetPresharedKey sets the preshared key for a peer.
|
// SetPresharedKey sets the preshared key for a peer.
|
||||||
@@ -92,14 +105,38 @@ func (c *WGUSPConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, updat
|
|||||||
}
|
}
|
||||||
|
|
||||||
cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly)
|
cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly)
|
||||||
return c.device.IpcSet(toWgUserspaceString(cfg))
|
if err := c.device.IpcSet(toWgUserspaceString(cfg)); err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Without updateOnly this creates the peer when it is absent, so the store has to
|
||||||
|
// know about it even though no allowed IP was configured.
|
||||||
|
if !updateOnly {
|
||||||
|
c.allowedIPs.ensure(parsedPeerKey)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set.
|
||||||
|
// It validates the endpoint before writing and records changes after a successful write.
|
||||||
func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error {
|
func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error {
|
||||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Everything that can fail is done before the device is touched, so a failure here
|
||||||
|
// cannot leave the device holding a peer that the activity recorder and the allowed
|
||||||
|
// IP store never learned about.
|
||||||
|
var addrPort netip.AddrPort
|
||||||
|
if endpoint != nil {
|
||||||
|
addr, err := netip.ParseAddr(endpoint.IP.String())
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("parse endpoint address: %w", err)
|
||||||
|
}
|
||||||
|
addrPort = netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port))
|
||||||
|
}
|
||||||
|
|
||||||
peer := wgtypes.PeerConfig{
|
peer := wgtypes.PeerConfig{
|
||||||
PublicKey: peerKeyParsed,
|
PublicKey: peerKeyParsed,
|
||||||
ReplaceAllowedIPs: false,
|
ReplaceAllowedIPs: false,
|
||||||
@@ -119,47 +156,27 @@ func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix,
|
|||||||
}
|
}
|
||||||
|
|
||||||
if endpoint != nil {
|
if endpoint != nil {
|
||||||
addr, err := netip.ParseAddr(endpoint.IP.String())
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to parse endpoint address: %w", err)
|
|
||||||
}
|
|
||||||
addrPort := netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port))
|
|
||||||
c.activityRecorder.UpsertAddress(peerKey, addrPort)
|
c.activityRecorder.UpsertAddress(peerKey, addrPort)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
c.allowedIPs.add(peerKeyParsed, allowedIps)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured.
|
||||||
|
// The UAPI cannot clear an endpoint in place, so the peer is removed and re-added with the
|
||||||
|
// allowed IPs it already had.
|
||||||
func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
|
func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
|
||||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("parse peer key: %w", err)
|
return fmt.Errorf("parse peer key: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
ipcStr, err := c.device.IpcGet()
|
allowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get IPC config: %w", err)
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse current status to get allowed IPs for the peer
|
|
||||||
stats, err := parseStatus(c.deviceName, ipcStr)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("parse IPC config: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var allowedIPs []net.IPNet
|
|
||||||
found := false
|
|
||||||
for _, peer := range stats.Peers {
|
|
||||||
if peer.PublicKey == peerKey {
|
|
||||||
allowedIPs = peer.AllowedIPs
|
|
||||||
found = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !found {
|
|
||||||
return fmt.Errorf("peer %s not found", peerKey)
|
|
||||||
}
|
|
||||||
|
|
||||||
// remove the peer from the WireGuard configuration
|
|
||||||
peer := wgtypes.PeerConfig{
|
peer := wgtypes.PeerConfig{
|
||||||
PublicKey: peerKeyParsed,
|
PublicKey: peerKeyParsed,
|
||||||
Remove: true,
|
Remove: true,
|
||||||
@@ -169,14 +186,13 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
|
|||||||
Peers: []wgtypes.PeerConfig{peer},
|
Peers: []wgtypes.PeerConfig{peer},
|
||||||
}
|
}
|
||||||
if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil {
|
if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil {
|
||||||
return fmt.Errorf("failed to remove peer: %s", ipcErr)
|
return fmt.Errorf("remove peer: %w", ipcErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build the peer config
|
|
||||||
peer = wgtypes.PeerConfig{
|
peer = wgtypes.PeerConfig{
|
||||||
PublicKey: peerKeyParsed,
|
PublicKey: peerKeyParsed,
|
||||||
ReplaceAllowedIPs: true,
|
ReplaceAllowedIPs: true,
|
||||||
AllowedIPs: allowedIPs,
|
AllowedIPs: prefixesToIPNets(allowedIPs),
|
||||||
}
|
}
|
||||||
|
|
||||||
config = wgtypes.Config{
|
config = wgtypes.Config{
|
||||||
@@ -184,12 +200,15 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
|
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
|
||||||
return fmt.Errorf("remove endpoint address: %w", err)
|
c.allowedIPs.forget(peerKeyParsed)
|
||||||
|
return fmt.Errorf("re-add peer without endpoint: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RemovePeer removes a peer, then clears its activity and allowed IP records.
|
||||||
|
// A failed device write leaves both records intact.
|
||||||
func (c *WGUSPConfigurer) RemovePeer(peerKey string) error {
|
func (c *WGUSPConfigurer) RemovePeer(peerKey string) error {
|
||||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -204,18 +223,17 @@ func (c *WGUSPConfigurer) RemovePeer(peerKey string) error {
|
|||||||
config := wgtypes.Config{
|
config := wgtypes.Config{
|
||||||
Peers: []wgtypes.PeerConfig{peer},
|
Peers: []wgtypes.PeerConfig{peer},
|
||||||
}
|
}
|
||||||
ipcErr := c.device.IpcSet(toWgUserspaceString(config))
|
if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil {
|
||||||
|
|
||||||
c.activityRecorder.Remove(peerKey)
|
|
||||||
return ipcErr
|
return ipcErr
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
c.activityRecorder.Remove(peerKey)
|
||||||
ipNet := net.IPNet{
|
c.allowedIPs.forget(peerKeyParsed)
|
||||||
IP: allowedIP.Addr().AsSlice(),
|
return nil
|
||||||
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op.
|
||||||
|
func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
||||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -224,79 +242,89 @@ func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) e
|
|||||||
PublicKey: peerKeyParsed,
|
PublicKey: peerKeyParsed,
|
||||||
UpdateOnly: true,
|
UpdateOnly: true,
|
||||||
ReplaceAllowedIPs: false,
|
ReplaceAllowedIPs: false,
|
||||||
AllowedIPs: []net.IPNet{ipNet},
|
AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}),
|
||||||
}
|
}
|
||||||
|
|
||||||
config := wgtypes.Config{
|
config := wgtypes.Config{
|
||||||
Peers: []wgtypes.PeerConfig{peer},
|
Peers: []wgtypes.PeerConfig{peer},
|
||||||
}
|
}
|
||||||
|
|
||||||
return c.device.IpcSet(toWgUserspaceString(config))
|
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
|
||||||
}
|
|
||||||
|
|
||||||
func (c *WGUSPConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
|
||||||
ipc, err := c.device.IpcGet()
|
|
||||||
if err != nil {
|
|
||||||
return err
|
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 {
|
||||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("parse peer key: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
hexKey := hex.EncodeToString(peerKeyParsed[:])
|
|
||||||
|
|
||||||
lines := strings.Split(ipc, "\n")
|
idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP))
|
||||||
|
if idx < 0 {
|
||||||
|
return ErrAllowedIPNotFound
|
||||||
|
}
|
||||||
|
newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1)
|
||||||
|
|
||||||
peer := wgtypes.PeerConfig{
|
peer := wgtypes.PeerConfig{
|
||||||
PublicKey: peerKeyParsed,
|
PublicKey: peerKeyParsed,
|
||||||
UpdateOnly: true,
|
UpdateOnly: true,
|
||||||
ReplaceAllowedIPs: true,
|
ReplaceAllowedIPs: true,
|
||||||
AllowedIPs: []net.IPNet{},
|
AllowedIPs: prefixesToIPNets(newAllowedIPs),
|
||||||
}
|
}
|
||||||
|
|
||||||
foundPeer := false
|
|
||||||
removedAllowedIP := false
|
|
||||||
ip := allowedIP.String()
|
|
||||||
|
|
||||||
for _, line := range lines {
|
|
||||||
line = strings.TrimSpace(line)
|
|
||||||
|
|
||||||
// If we're within the details of the found peer and encounter another public key,
|
|
||||||
// this means we're starting another peer's details. So, reset the flag.
|
|
||||||
if strings.HasPrefix(line, "public_key=") && foundPeer {
|
|
||||||
foundPeer = false
|
|
||||||
}
|
|
||||||
|
|
||||||
// Identify the peer with the specific public key
|
|
||||||
if line == fmt.Sprintf("public_key=%s", hexKey) {
|
|
||||||
foundPeer = true
|
|
||||||
}
|
|
||||||
|
|
||||||
// If we're within the details of the found peer and find the specific allowed IP, skip this line
|
|
||||||
if foundPeer && line == "allowed_ip="+ip {
|
|
||||||
removedAllowedIP = true
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
// Append the line to the output string
|
|
||||||
if foundPeer && strings.HasPrefix(line, "allowed_ip=") {
|
|
||||||
allowedIPStr := strings.TrimPrefix(line, "allowed_ip=")
|
|
||||||
_, ipNet, err := net.ParseCIDR(allowedIPStr)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
peer.AllowedIPs = append(peer.AllowedIPs, *ipNet)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !removedAllowedIP {
|
|
||||||
return ErrAllowedIPNotFound
|
|
||||||
}
|
|
||||||
config := wgtypes.Config{
|
config := wgtypes.Config{
|
||||||
Peers: []wgtypes.PeerConfig{peer},
|
Peers: []wgtypes.PeerConfig{peer},
|
||||||
}
|
}
|
||||||
return c.device.IpcSet(toWgUserspaceString(config))
|
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
|
||||||
|
return fmt.Errorf("remove allowed IP %s: %w", allowedIP, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
c.allowedIPs.set(peerKeyParsed, newAllowedIPs)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device
|
||||||
|
// only for a peer the store has not seen. Reading them back means dumping and parsing the
|
||||||
|
// whole device configuration, and this runs on every relay and ICE transition.
|
||||||
|
func (c *WGUSPConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) {
|
||||||
|
if prefixes, ok := c.allowedIPs.get(peerKey); ok {
|
||||||
|
return prefixes, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
ipcStr, err := c.device.IpcGet()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("get IPC config: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
stats, err := parseStatus(c.deviceName, ipcStr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("parse IPC config: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseStatus reports keys in their textual form, so the comparison needs it once.
|
||||||
|
wanted := peerKey.String()
|
||||||
|
for _, peer := range stats.Peers {
|
||||||
|
if peer.PublicKey != wanted {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
prefixes := ipNetsToPrefixes(peer.AllowedIPs)
|
||||||
|
c.allowedIPs.set(peerKey, prefixes)
|
||||||
|
return prefixes, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, ErrPeerNotFound
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *WGUSPConfigurer) FullStats() (*Stats, error) {
|
func (c *WGUSPConfigurer) FullStats() (*Stats, error) {
|
||||||
|
|||||||
@@ -0,0 +1,318 @@
|
|||||||
|
package configurer
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
wgconn "golang.zx2c4.com/wireguard/conn"
|
||||||
|
wgdevice "golang.zx2c4.com/wireguard/device"
|
||||||
|
"golang.zx2c4.com/wireguard/tun/tuntest"
|
||||||
|
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/iface/bind"
|
||||||
|
)
|
||||||
|
|
||||||
|
// newTestUSPConfigurer builds a configurer over a real wireguard-go device backed by an
|
||||||
|
// in-memory TUN. The device stays down, so no socket is opened and no privileges are needed.
|
||||||
|
func newTestUSPConfigurer(t *testing.T) *WGUSPConfigurer {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
tun := tuntest.NewChannelTUN()
|
||||||
|
dev := wgdevice.NewDevice(tun.TUN(), wgconn.NewDefaultBind(), wgdevice.NewLogger(wgdevice.LogLevelSilent, ""))
|
||||||
|
t.Cleanup(dev.Close)
|
||||||
|
|
||||||
|
c := NewUSPConfigurerNoUAPI(dev, "wgtest0", bind.NewActivityRecorder())
|
||||||
|
|
||||||
|
key, err := wgtypes.GeneratePrivateKey()
|
||||||
|
require.NoError(t, err, "generate device private key")
|
||||||
|
require.NoError(t, c.ConfigureInterface(key.String(), 0), "configure test device")
|
||||||
|
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// seedPeers adds count peers, each with a /32 overlay address, and returns their public keys.
|
||||||
|
func seedPeers(t *testing.T, c *WGUSPConfigurer, count int) []string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
keys := make([]string, 0, count)
|
||||||
|
for i := 0; i < count; i++ {
|
||||||
|
priv, err := wgtypes.GeneratePrivateKey()
|
||||||
|
require.NoError(t, err, "generate peer private key")
|
||||||
|
pub := priv.PublicKey().String()
|
||||||
|
|
||||||
|
addr := netip.PrefixFrom(netip.AddrFrom4([4]byte{100, 64, byte(i >> 8), byte(i)}), 32)
|
||||||
|
require.NoError(t, c.UpdatePeer(pub, []netip.Prefix{addr}, 25*time.Second, nil, nil), "add peer")
|
||||||
|
keys = append(keys, pub)
|
||||||
|
}
|
||||||
|
return keys
|
||||||
|
}
|
||||||
|
|
||||||
|
func peerAllowedIPs(t *testing.T, c *WGUSPConfigurer, peerKey string) []string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
stats, err := c.FullStats()
|
||||||
|
require.NoError(t, err, "read device stats")
|
||||||
|
|
||||||
|
for _, p := range stats.Peers {
|
||||||
|
if p.PublicKey != peerKey {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
got := make([]string, 0, len(p.AllowedIPs))
|
||||||
|
for _, ipNet := range p.AllowedIPs {
|
||||||
|
got = append(got, ipNet.String())
|
||||||
|
}
|
||||||
|
return got
|
||||||
|
}
|
||||||
|
t.Fatalf("peer %s not found on device", peerKey)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRemoveEndpointAddressPreservesRoutedAllowedIPs covers the prefixes the route manager
|
||||||
|
// attaches to a routing peer through AddAllowedIP. Those are not known to the peer.Conn that
|
||||||
|
// triggers the endpoint removal, so dropping them here would silently blackhole every route
|
||||||
|
// behind that peer on each relay or ICE disconnect.
|
||||||
|
func TestRemoveEndpointAddressPreservesRoutedAllowedIPs(t *testing.T) {
|
||||||
|
c := newTestUSPConfigurer(t)
|
||||||
|
peerKey := seedPeers(t, c, 3)[1]
|
||||||
|
|
||||||
|
routed := []netip.Prefix{
|
||||||
|
netip.MustParsePrefix("10.20.0.0/16"),
|
||||||
|
netip.MustParsePrefix("192.168.7.0/24"),
|
||||||
|
}
|
||||||
|
for _, prefix := range routed {
|
||||||
|
require.NoError(t, c.AddAllowedIP(peerKey, prefix), "add routed prefix")
|
||||||
|
}
|
||||||
|
|
||||||
|
before := peerAllowedIPs(t, c, peerKey)
|
||||||
|
require.Len(t, before, 3, "peer should hold its overlay address plus both routed prefixes")
|
||||||
|
|
||||||
|
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
|
||||||
|
|
||||||
|
assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey),
|
||||||
|
"allowed IPs must survive the endpoint removal unchanged")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRemoveEndpointAddressDoesNotScaleWithPeerCount is the regression guard for the actual
|
||||||
|
// defect: clearing one peer's endpoint used to dump and parse the whole device, so its cost
|
||||||
|
// grew with the size of the network map. On a routing peer with thousands of peers that dump
|
||||||
|
// runs on every relay and ICE transition, under the interface lock.
|
||||||
|
func TestRemoveEndpointAddressDoesNotScaleWithPeerCount(t *testing.T) {
|
||||||
|
measure := func(peerCount int) float64 {
|
||||||
|
c := newTestUSPConfigurer(t)
|
||||||
|
peerKey := seedPeers(t, c, peerCount)[peerCount/2]
|
||||||
|
|
||||||
|
return testing.AllocsPerRun(5, func() {
|
||||||
|
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
small := measure(64)
|
||||||
|
large := measure(1024)
|
||||||
|
|
||||||
|
assert.Less(t, large, small*2,
|
||||||
|
"clearing one endpoint allocated %.0f objects with 1024 peers against %.0f with 64: the cost still scales with the peer count",
|
||||||
|
large, small)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRemoveEndpointAddressFallsBackToDevice covers a peer the store never saw, which is what
|
||||||
|
// an out-of-band reconfiguration of the device leaves behind. The device stays the source of
|
||||||
|
// truth in that case, so the allowed IPs must still be preserved.
|
||||||
|
func TestRemoveEndpointAddressFallsBackToDevice(t *testing.T) {
|
||||||
|
c := newTestUSPConfigurer(t)
|
||||||
|
peerKey := seedPeers(t, c, 3)[1]
|
||||||
|
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix")
|
||||||
|
|
||||||
|
before := peerAllowedIPs(t, c, peerKey)
|
||||||
|
c.allowedIPs.reset()
|
||||||
|
|
||||||
|
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
|
||||||
|
|
||||||
|
assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey),
|
||||||
|
"allowed IPs recovered from the device must be preserved")
|
||||||
|
|
||||||
|
recovered, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
|
||||||
|
assert.True(t, ok, "the fallback must seed the store so the next call skips the device dump")
|
||||||
|
assert.Len(t, recovered, 2, "seeded prefixes")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoveAllowedIPKeepsTheOtherPrefixes(t *testing.T) {
|
||||||
|
c := newTestUSPConfigurer(t)
|
||||||
|
peerKey := seedPeers(t, c, 3)[0]
|
||||||
|
routed := netip.MustParsePrefix("10.20.0.0/16")
|
||||||
|
require.NoError(t, c.AddAllowedIP(peerKey, routed), "add routed prefix")
|
||||||
|
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("192.168.7.0/24")), "add routed prefix")
|
||||||
|
|
||||||
|
require.NoError(t, c.RemoveAllowedIP(peerKey, routed), "remove routed prefix")
|
||||||
|
|
||||||
|
assert.ElementsMatch(t, []string{"100.64.0.0/32", "192.168.7.0/24"}, peerAllowedIPs(t, c, peerKey),
|
||||||
|
"only the removed prefix should be gone")
|
||||||
|
|
||||||
|
assert.ErrorIs(t, c.RemoveAllowedIP(peerKey, routed), ErrAllowedIPNotFound,
|
||||||
|
"removing a prefix that is no longer configured must be reported")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt covers the lazy connection window documented
|
||||||
|
// in #6863: AddAllowedIP is update-only, a silent no-op when the peer is absent, so it must not
|
||||||
|
// leave the store claiming prefixes the device never took. RemoveEndpointAddress re-adds a peer
|
||||||
|
// without update-only, so a phantom entry would create a peer the device had dropped, and a
|
||||||
|
// created peer would steal those allowed IPs from whichever peer legitimately holds them.
|
||||||
|
func TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt(t *testing.T) {
|
||||||
|
c := newTestUSPConfigurer(t)
|
||||||
|
seedPeers(t, c, 2)
|
||||||
|
|
||||||
|
priv, err := wgtypes.GeneratePrivateKey()
|
||||||
|
require.NoError(t, err, "generate peer private key")
|
||||||
|
absent := priv.PublicKey().String()
|
||||||
|
|
||||||
|
require.NoError(t, c.AddAllowedIP(absent, netip.MustParsePrefix("10.20.0.0/16")),
|
||||||
|
"update-only add on an absent peer is a silent no-op")
|
||||||
|
|
||||||
|
stats, err := c.FullStats()
|
||||||
|
require.NoError(t, err, "read device stats")
|
||||||
|
require.Len(t, stats.Peers, 2, "the absent peer must not have been created by AddAllowedIP")
|
||||||
|
|
||||||
|
assert.ErrorIs(t, c.RemoveEndpointAddress(absent), ErrPeerNotFound,
|
||||||
|
"clearing the endpoint of a peer the device does not have must fail")
|
||||||
|
|
||||||
|
stats, err = c.FullStats()
|
||||||
|
require.NoError(t, err, "read device stats")
|
||||||
|
assert.Len(t, stats.Peers, 2, "no peer may be created while clearing an endpoint")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer covers WireGuard's rule that an
|
||||||
|
// allowed IP belongs to exactly one peer: configuring a prefix on a peer takes it away from
|
||||||
|
// whichever peer held it before. UpdatePeer relies on that rule rather than removing the prefix
|
||||||
|
// from the previous holder itself, so a prefix handed over between peers must not come back.
|
||||||
|
func TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer(t *testing.T) {
|
||||||
|
c := newTestUSPConfigurer(t)
|
||||||
|
keys := seedPeers(t, c, 2)
|
||||||
|
peerA, peerB := keys[0], keys[1]
|
||||||
|
routed := netip.MustParsePrefix("10.20.0.0/16")
|
||||||
|
|
||||||
|
require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A")
|
||||||
|
require.Contains(t, peerAllowedIPs(t, c, peerA), routed.String(), "A must hold the prefix")
|
||||||
|
|
||||||
|
// The route moves to B. The device takes it away from A on its own.
|
||||||
|
require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B")
|
||||||
|
require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix")
|
||||||
|
require.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(), "the device must have taken it from A")
|
||||||
|
|
||||||
|
require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint")
|
||||||
|
|
||||||
|
assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(),
|
||||||
|
"clearing A's endpoint must not take the prefix back from B")
|
||||||
|
assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(),
|
||||||
|
"B must still hold the prefix")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestPresharedKeyCreatedPeerTakesPartInPrefixHandover covers a peer created by a preshared
|
||||||
|
// key write rather than by a peer update. Rosenpass applies a peer's first key without
|
||||||
|
// updateOnly, which creates the peer on the device, so a store that ignored that operation
|
||||||
|
// would treat the peer as unknown and would not account for a prefix later handed over to it.
|
||||||
|
func TestPresharedKeyCreatedPeerTakesPartInPrefixHandover(t *testing.T) {
|
||||||
|
c := newTestUSPConfigurer(t)
|
||||||
|
peerA := seedPeers(t, c, 1)[0]
|
||||||
|
routed := netip.MustParsePrefix("10.20.0.0/16")
|
||||||
|
require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A")
|
||||||
|
|
||||||
|
priv, err := wgtypes.GeneratePrivateKey()
|
||||||
|
require.NoError(t, err, "generate peer private key")
|
||||||
|
peerB := priv.PublicKey().String()
|
||||||
|
|
||||||
|
psk, err := wgtypes.GenerateKey()
|
||||||
|
require.NoError(t, err, "generate preshared key")
|
||||||
|
require.NoError(t, c.SetPresharedKey(peerB, psk, false), "a first key creates the peer")
|
||||||
|
|
||||||
|
require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B")
|
||||||
|
require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix")
|
||||||
|
|
||||||
|
require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint")
|
||||||
|
|
||||||
|
assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(),
|
||||||
|
"clearing A's endpoint must not take the prefix back from B")
|
||||||
|
assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must still hold the prefix")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice is the end to end form of the
|
||||||
|
// conversion: a v4-mapped prefix must not reach the device as a zero length allowed IP,
|
||||||
|
// which would route every v4 address to that peer.
|
||||||
|
func TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice(t *testing.T) {
|
||||||
|
c := newTestUSPConfigurer(t)
|
||||||
|
|
||||||
|
priv, err := wgtypes.GeneratePrivateKey()
|
||||||
|
require.NoError(t, err, "generate peer private key")
|
||||||
|
peerKey := priv.PublicKey().String()
|
||||||
|
|
||||||
|
mapped := netip.MustParsePrefix("::ffff:10.1.2.3/112")
|
||||||
|
require.NoError(t, c.UpdatePeer(peerKey, []netip.Prefix{mapped}, 25*time.Second, nil, nil), "add peer")
|
||||||
|
|
||||||
|
onDevice := peerAllowedIPs(t, c, peerKey)
|
||||||
|
assert.NotContains(t, onDevice, "0.0.0.0/0", "the device must not be given a catch-all allowed IP")
|
||||||
|
assert.Equal(t, []string{"10.1.0.0/16"}, onDevice, "the device holds the normalized prefix")
|
||||||
|
|
||||||
|
recorded, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
|
||||||
|
require.True(t, ok, "the peer must be recorded")
|
||||||
|
require.Len(t, recorded, 1, "one prefix recorded")
|
||||||
|
assert.Equal(t, onDevice[0], recorded[0].String(), "device and store must agree")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestUpdatePeerWithAnUnusableEndpointTouchesNothing pins the ordering: the endpoint is
|
||||||
|
// parsed before the device is configured, so a failure cannot leave the device holding a
|
||||||
|
// peer that the store never learned about, with the prefix handover skipped along with it.
|
||||||
|
func TestUpdatePeerWithAnUnusableEndpointTouchesNothing(t *testing.T) {
|
||||||
|
c := newTestUSPConfigurer(t)
|
||||||
|
seedPeers(t, c, 2)
|
||||||
|
|
||||||
|
priv, err := wgtypes.GeneratePrivateKey()
|
||||||
|
require.NoError(t, err, "generate peer private key")
|
||||||
|
peerKey := priv.PublicKey().String()
|
||||||
|
|
||||||
|
// A three byte address has no textual form netip can parse back.
|
||||||
|
endpoint := &net.UDPAddr{IP: net.IP{1, 2, 3}, Port: 51820}
|
||||||
|
require.Error(t, c.UpdatePeer(peerKey, []netip.Prefix{netip.MustParsePrefix("10.30.0.0/16")},
|
||||||
|
25*time.Second, endpoint, nil), "an unusable endpoint must fail the update")
|
||||||
|
|
||||||
|
stats, err := c.FullStats()
|
||||||
|
require.NoError(t, err, "read device stats")
|
||||||
|
assert.Len(t, stats.Peers, 2, "the peer must not have reached the device")
|
||||||
|
|
||||||
|
_, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
|
||||||
|
assert.False(t, ok, "the peer must not have been recorded either")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses covers a removal that never reached the
|
||||||
|
// device. A single peer removal is one write, so a failure leaves the peer on the device
|
||||||
|
// exactly as it was, and the record still describes it; dropping it would only force the
|
||||||
|
// next caller to read the whole device back for an answer it already had.
|
||||||
|
func TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses(t *testing.T) {
|
||||||
|
c := newTestUSPConfigurer(t)
|
||||||
|
peerKey := seedPeers(t, c, 1)[0]
|
||||||
|
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix")
|
||||||
|
|
||||||
|
before, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
|
||||||
|
require.True(t, ok, "the peer must be recorded before the removal")
|
||||||
|
require.Len(t, before, 2, "overlay address plus routed prefix")
|
||||||
|
|
||||||
|
// A closed device refuses every write, which is the shape of any failed removal.
|
||||||
|
c.device.Close()
|
||||||
|
|
||||||
|
require.Error(t, c.RemovePeer(peerKey), "the removal must report the failure")
|
||||||
|
|
||||||
|
after, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
|
||||||
|
require.True(t, ok, "a peer still on the device must stay recorded")
|
||||||
|
assert.Equal(t, before, after, "the record must describe the peer the device kept")
|
||||||
|
}
|
||||||
|
|
||||||
|
// mustParseKey turns the textual key the configurer API takes into the form the store
|
||||||
|
// keys on.
|
||||||
|
func mustParseKey(t *testing.T, key string) wgtypes.Key {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
parsed, err := wgtypes.ParseKey(key)
|
||||||
|
require.NoError(t, err, "parse peer key")
|
||||||
|
return parsed
|
||||||
|
}
|
||||||
@@ -51,7 +51,6 @@ func ValidateMTU(mtu uint16) error {
|
|||||||
|
|
||||||
type wgProxyFactory interface {
|
type wgProxyFactory interface {
|
||||||
GetProxy() wgproxy.Proxy
|
GetProxy() wgproxy.Proxy
|
||||||
GetProxyPort() uint16
|
|
||||||
Free() error
|
Free() error
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -81,12 +80,6 @@ func (w *WGIface) GetProxy() wgproxy.Proxy {
|
|||||||
return w.wgProxyFactory.GetProxy()
|
return w.wgProxyFactory.GetProxy()
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetProxyPort returns the proxy port used by the WireGuard proxy.
|
|
||||||
// Returns 0 if no proxy port is used (e.g., for userspace WireGuard).
|
|
||||||
func (w *WGIface) GetProxyPort() uint16 {
|
|
||||||
return w.wgProxyFactory.GetProxyPort()
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetBind returns the EndpointManager userspace bind mode.
|
// GetBind returns the EndpointManager userspace bind mode.
|
||||||
func (w *WGIface) GetBind() device.EndpointManager {
|
func (w *WGIface) GetBind() device.EndpointManager {
|
||||||
w.mu.Lock()
|
w.mu.Lock()
|
||||||
|
|||||||
@@ -52,7 +52,6 @@ func (f *fakeTunDevice) Close() error {
|
|||||||
type fakeProxyFactory struct{}
|
type fakeProxyFactory struct{}
|
||||||
|
|
||||||
func (fakeProxyFactory) GetProxy() wgproxy.Proxy { return nil }
|
func (fakeProxyFactory) GetProxy() wgproxy.Proxy { return nil }
|
||||||
func (fakeProxyFactory) GetProxyPort() uint16 { return 0 }
|
|
||||||
func (fakeProxyFactory) Free() error { return nil }
|
func (fakeProxyFactory) Free() error { return nil }
|
||||||
|
|
||||||
// TestWGIface_CloseReleasesMutexBeforeTunClose guards against a deadlock
|
// TestWGIface_CloseReleasesMutexBeforeTunClose guards against a deadlock
|
||||||
|
|||||||
@@ -6,27 +6,14 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
"github.com/netbirdio/netbird/client/internal/wincmd"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (w *WGIface) Destroy() error {
|
func (w *WGIface) Destroy() error {
|
||||||
netshCmd := GetSystem32Command("netsh")
|
netshCmd := wincmd.System32("netsh")
|
||||||
out, err := exec.Command(netshCmd, "interface", "set", "interface", w.Name(), "admin=disable").CombinedOutput()
|
out, err := exec.Command(netshCmd, "interface", "set", "interface", w.Name(), "admin=disable").CombinedOutput()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to remove interface %s: %w - %s", w.Name(), err, out)
|
return fmt.Errorf("failed to remove interface %s: %w - %s", w.Name(), err, out)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetSystem32Command checks if a command can be found in the system path and returns it. In case it can't find it
|
|
||||||
// in the path it will return the full path of a command assuming C:\windows\system32 as the base path.
|
|
||||||
func GetSystem32Command(command string) string {
|
|
||||||
_, err := exec.LookPath(command)
|
|
||||||
if err == nil {
|
|
||||||
return command
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Tracef("Command %s not found in PATH, using C:\\windows\\system32\\%s.exe path", command, command)
|
|
||||||
|
|
||||||
return "C:\\windows\\system32\\" + command + ".exe"
|
|
||||||
}
|
|
||||||
|
|||||||
+38
-42
@@ -40,14 +40,18 @@ func init() {
|
|||||||
peerPubKey = peerPrivateKey.PublicKey().String()
|
peerPubKey = peerPrivateKey.PublicKey().String()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// testIFaceBlackList mirrors the prefixes profilemanager.DefaultInterfaceBlacklist
|
||||||
|
// carries for the overlay interface. These tests create their own utun device, and
|
||||||
|
// stdnet's filter probes with wgctrl every interface it is not told to skip, which
|
||||||
|
// on a userspace WireGuard platform reaches the UAPI socket of this same process.
|
||||||
|
// Declared here rather than imported because profilemanager imports this package.
|
||||||
|
var testIFaceBlackList = []string{"wt", "utun", "tun0"}
|
||||||
|
|
||||||
func TestWGIface_UpdateAddr(t *testing.T) {
|
func TestWGIface_UpdateAddr(t *testing.T) {
|
||||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
||||||
addr := "100.64.0.1/8"
|
addr := "100.64.0.1/8"
|
||||||
wgPort := 33100
|
wgPort := 33100
|
||||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
opts := WGIFaceOpts{
|
opts := WGIFaceOpts{
|
||||||
IFaceName: ifaceName,
|
IFaceName: ifaceName,
|
||||||
@@ -127,10 +131,7 @@ func getIfaceAddrs(ifaceName string) ([]net.Addr, error) {
|
|||||||
func Test_CreateInterface(t *testing.T) {
|
func Test_CreateInterface(t *testing.T) {
|
||||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1)
|
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1)
|
||||||
wgIP := "10.99.99.1/32"
|
wgIP := "10.99.99.1/32"
|
||||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
opts := WGIFaceOpts{
|
opts := WGIFaceOpts{
|
||||||
IFaceName: ifaceName,
|
IFaceName: ifaceName,
|
||||||
Address: wgaddr.MustParseWGAddress(wgIP),
|
Address: wgaddr.MustParseWGAddress(wgIP),
|
||||||
@@ -170,10 +171,7 @@ func Test_Close(t *testing.T) {
|
|||||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
|
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
|
||||||
wgIP := "10.99.99.2/32"
|
wgIP := "10.99.99.2/32"
|
||||||
wgPort := 33100
|
wgPort := 33100
|
||||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
opts := WGIFaceOpts{
|
opts := WGIFaceOpts{
|
||||||
IFaceName: ifaceName,
|
IFaceName: ifaceName,
|
||||||
@@ -215,10 +213,7 @@ func TestRecreation(t *testing.T) {
|
|||||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
|
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
|
||||||
wgIP := "10.99.99.2/32"
|
wgIP := "10.99.99.2/32"
|
||||||
wgPort := 33100
|
wgPort := 33100
|
||||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
opts := WGIFaceOpts{
|
opts := WGIFaceOpts{
|
||||||
IFaceName: ifaceName,
|
IFaceName: ifaceName,
|
||||||
@@ -288,10 +283,7 @@ func Test_ConfigureInterface(t *testing.T) {
|
|||||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3)
|
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3)
|
||||||
wgIP := "10.99.99.5/30"
|
wgIP := "10.99.99.5/30"
|
||||||
wgPort := 33100
|
wgPort := 33100
|
||||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
opts := WGIFaceOpts{
|
opts := WGIFaceOpts{
|
||||||
IFaceName: ifaceName,
|
IFaceName: ifaceName,
|
||||||
Address: wgaddr.MustParseWGAddress(wgIP),
|
Address: wgaddr.MustParseWGAddress(wgIP),
|
||||||
@@ -343,10 +335,7 @@ func Test_ConfigureInterface(t *testing.T) {
|
|||||||
func Test_UpdatePeer(t *testing.T) {
|
func Test_UpdatePeer(t *testing.T) {
|
||||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
||||||
wgIP := "10.99.99.9/30"
|
wgIP := "10.99.99.9/30"
|
||||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
opts := WGIFaceOpts{
|
opts := WGIFaceOpts{
|
||||||
IFaceName: ifaceName,
|
IFaceName: ifaceName,
|
||||||
@@ -413,10 +402,7 @@ func Test_UpdatePeer(t *testing.T) {
|
|||||||
func Test_RemovePeer(t *testing.T) {
|
func Test_RemovePeer(t *testing.T) {
|
||||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
||||||
wgIP := "10.99.99.13/30"
|
wgIP := "10.99.99.13/30"
|
||||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
opts := WGIFaceOpts{
|
opts := WGIFaceOpts{
|
||||||
IFaceName: ifaceName,
|
IFaceName: ifaceName,
|
||||||
@@ -477,10 +463,7 @@ func Test_ConnectPeers(t *testing.T) {
|
|||||||
peer2wgPort := 33200
|
peer2wgPort := 33200
|
||||||
|
|
||||||
keepAlive := 1 * time.Second
|
keepAlive := 1 * time.Second
|
||||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
guid := fmt.Sprintf("{%s}", uuid.New().String())
|
guid := fmt.Sprintf("{%s}", uuid.New().String())
|
||||||
device.CustomWindowsGUIDString = strings.ToLower(guid)
|
device.CustomWindowsGUIDString = strings.ToLower(guid)
|
||||||
@@ -516,10 +499,7 @@ func Test_ConnectPeers(t *testing.T) {
|
|||||||
guid = fmt.Sprintf("{%s}", uuid.New().String())
|
guid = fmt.Sprintf("{%s}", uuid.New().String())
|
||||||
device.CustomWindowsGUIDString = strings.ToLower(guid)
|
device.CustomWindowsGUIDString = strings.ToLower(guid)
|
||||||
|
|
||||||
newNet, err = stdnet.NewNet(context.Background(), nil)
|
newNet = stdnet.NewNet(context.Background(), testIFaceBlackList, nil)
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
optsPeer2 := WGIFaceOpts{
|
optsPeer2 := WGIFaceOpts{
|
||||||
IFaceName: peer2ifaceName,
|
IFaceName: peer2ifaceName,
|
||||||
@@ -568,11 +548,14 @@ func Test_ConnectPeers(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
// The peers use userspace WireGuard (stdnet transport). A tight busy-loop
|
// On Linux with the kernel module both peers are kernel devices, elsewhere
|
||||||
// here starves the wireguard-go goroutines that process the handshake, so
|
// they run on wireguard-go. A tight busy-loop here would starve the
|
||||||
// poll on a ticker instead and yield the CPU between checks. WireGuard also
|
// wireguard-go goroutines that process the handshake, so poll on a ticker
|
||||||
// only retries a lost handshake initiation every REKEY_TIMEOUT (5s), which
|
// instead and yield the CPU between checks. WireGuard also only retries a
|
||||||
// is why the overall wait can occasionally stretch to tens of seconds.
|
// lost handshake initiation every REKEY_TIMEOUT (5s), which is why the
|
||||||
|
// overall wait can occasionally stretch to tens of seconds. Each side sends
|
||||||
|
// its first initiation when its peer is configured, and the first one leaves
|
||||||
|
// before the other device knows the peer, so that one is always wasted.
|
||||||
timeout := 30 * time.Second
|
timeout := 30 * time.Second
|
||||||
timeoutChannel := time.After(timeout)
|
timeoutChannel := time.After(timeout)
|
||||||
ticker := time.NewTicker(500 * time.Millisecond)
|
ticker := time.NewTicker(500 * time.Millisecond)
|
||||||
@@ -590,13 +573,26 @@ func Test_ConnectPeers(t *testing.T) {
|
|||||||
|
|
||||||
select {
|
select {
|
||||||
case <-timeoutChannel:
|
case <-timeoutChannel:
|
||||||
t.Fatalf("waiting for peer handshake timeout after %s", timeout.String())
|
// The counters tell whether initiations were sent at all, whether they
|
||||||
|
// arrived, and whether only one direction is working.
|
||||||
|
t.Fatalf("waiting for peer handshake timeout after %s\n%s\n%s", timeout.String(),
|
||||||
|
describePeer(peer1ifaceName, peer2Key.PublicKey().String()),
|
||||||
|
describePeer(peer2ifaceName, peer1Key.PublicKey().String()))
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func describePeer(ifaceName, peerPubKey string) string {
|
||||||
|
peer, err := getPeer(ifaceName, peerPubKey)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Sprintf("%s: peer %s: %v", ifaceName, peerPubKey, err)
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%s: peer %s endpoint=%v tx=%d rx=%d last_handshake=%v",
|
||||||
|
ifaceName, peerPubKey, peer.Endpoint, peer.TransmitBytes, peer.ReceiveBytes, peer.LastHandshakeTime)
|
||||||
|
}
|
||||||
|
|
||||||
func getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) {
|
func getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) {
|
||||||
wg, err := wgctrl.New()
|
wg, err := wgctrl.New()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -200,10 +200,7 @@ func (m *SingleSocketUDPMux) updateLocalAddresses() {
|
|||||||
}
|
}
|
||||||
if len(networks) > 0 {
|
if len(networks) > 0 {
|
||||||
if m.params.Net == nil {
|
if m.params.Net == nil {
|
||||||
var err error
|
m.params.Net = stdnet.NewNet(context.Background(), nil, nil)
|
||||||
if m.params.Net, err = stdnet.NewNet(context.Background(), nil); err != nil {
|
|
||||||
m.params.Logger.Errorf("failed to get create network: %v", err)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
ips, err := localInterfaces(m.params.Net, m.params.InterfaceFilter, nil, networks, true)
|
ips, err := localInterfaces(m.params.Net, m.params.InterfaceFilter, nil, networks, true)
|
||||||
|
|||||||
@@ -1,32 +0,0 @@
|
|||||||
package ebpf
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"net"
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
portRangeStart = 3128
|
|
||||||
portRangeEnd = portRangeStart + 100
|
|
||||||
)
|
|
||||||
|
|
||||||
type portLookup struct {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (pl portLookup) searchFreePort() (int, error) {
|
|
||||||
for i := portRangeStart; i <= portRangeEnd; i++ {
|
|
||||||
if pl.tryToBind(i) == nil {
|
|
||||||
return i, nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return 0, fmt.Errorf("failed to bind free port for eBPF proxy")
|
|
||||||
}
|
|
||||||
|
|
||||||
func (pl portLookup) tryToBind(port int) error {
|
|
||||||
l, err := net.ListenPacket("udp", fmt.Sprintf(":%d", port))
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
_ = l.Close()
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,45 +0,0 @@
|
|||||||
package ebpf
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"net"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func Test_portLookup_searchFreePort(t *testing.T) {
|
|
||||||
pl := portLookup{}
|
|
||||||
_, err := pl.searchFreePort()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_portLookup_on_allocated(t *testing.T) {
|
|
||||||
pl := portLookup{}
|
|
||||||
|
|
||||||
portRangeStart = 4128
|
|
||||||
portRangeEnd = portRangeStart + 100
|
|
||||||
|
|
||||||
allocatedPort, err := allocatePort(portRangeStart)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
defer allocatedPort.Close()
|
|
||||||
|
|
||||||
fp, err := pl.searchFreePort()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if fp != (portRangeStart + 1) {
|
|
||||||
t.Errorf("invalid free port, expected: %d, got: %d", portRangeStart+1, fp)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func allocatePort(port int) (net.PacketConn, error) {
|
|
||||||
c, err := net.ListenPacket("udp", fmt.Sprintf(":%d", port))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return c, err
|
|
||||||
}
|
|
||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -8,11 +8,13 @@ import (
|
|||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/iface/wgproxy/ebpf"
|
"github.com/netbirdio/netbird/client/iface/wgproxy/loopback"
|
||||||
udpProxy "github.com/netbirdio/netbird/client/iface/wgproxy/udp"
|
udpProxy "github.com/netbirdio/netbird/client/iface/wgproxy/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
envDisableKernelWGProxy = "NB_DISABLE_KERNEL_WG_PROXY"
|
||||||
|
// envDisableEBPFWGProxy is a deprecated alias for envDisableKernelWGProxy.
|
||||||
envDisableEBPFWGProxy = "NB_DISABLE_EBPF_WG_PROXY"
|
envDisableEBPFWGProxy = "NB_DISABLE_EBPF_WG_PROXY"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -20,7 +22,7 @@ type KernelFactory struct {
|
|||||||
wgPort int
|
wgPort int
|
||||||
mtu uint16
|
mtu uint16
|
||||||
|
|
||||||
ebpfProxy *ebpf.WGEBPFProxy
|
loopbackProxy *loopback.Proxy
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory {
|
func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory {
|
||||||
@@ -29,55 +31,56 @@ func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory {
|
|||||||
mtu: mtu,
|
mtu: mtu,
|
||||||
}
|
}
|
||||||
|
|
||||||
if isEBPFDisabled() {
|
if isKernelProxyDisabled() {
|
||||||
log.Infof("WireGuard Proxy Factory will produce UDP proxy")
|
log.Infof("WireGuard Proxy Factory will produce UDP proxy")
|
||||||
log.Infof("eBPF WireGuard proxy is disabled via %s environment variable", envDisableEBPFWGProxy)
|
|
||||||
return f
|
return f
|
||||||
}
|
}
|
||||||
|
|
||||||
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, mtu)
|
loopbackProxy := loopback.NewProxy(wgPort, mtu)
|
||||||
if err := ebpfProxy.Listen(); err != nil {
|
if err := loopbackProxy.Listen(); err != nil {
|
||||||
log.Infof("WireGuard Proxy Factory will produce UDP proxy")
|
log.Infof("WireGuard Proxy Factory will produce UDP proxy")
|
||||||
log.Warnf("failed to initialize ebpf proxy, fallback to user space proxy: %s", err)
|
log.Warnf("failed to initialize loopback proxy, fallback to user space proxy: %s", err)
|
||||||
return f
|
return f
|
||||||
}
|
}
|
||||||
log.Infof("WireGuard Proxy Factory will produce eBPF proxy")
|
log.Infof("WireGuard Proxy Factory will produce loopback proxy")
|
||||||
f.ebpfProxy = ebpfProxy
|
f.loopbackProxy = loopbackProxy
|
||||||
return f
|
return f
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *KernelFactory) GetProxy() Proxy {
|
func (w *KernelFactory) GetProxy() Proxy {
|
||||||
if w.ebpfProxy == nil {
|
if w.loopbackProxy == nil {
|
||||||
return udpProxy.NewWGUDPProxy(w.wgPort, w.mtu)
|
return udpProxy.NewWGUDPProxy(w.wgPort, w.mtu)
|
||||||
}
|
}
|
||||||
|
|
||||||
return ebpf.NewProxyWrapper(w.ebpfProxy)
|
return loopback.NewProxyWrapper(w.loopbackProxy)
|
||||||
}
|
|
||||||
|
|
||||||
// GetProxyPort returns the eBPF proxy port, or 0 if eBPF is not active.
|
|
||||||
func (w *KernelFactory) GetProxyPort() uint16 {
|
|
||||||
if w.ebpfProxy == nil {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
return w.ebpfProxy.GetProxyPort()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *KernelFactory) Free() error {
|
func (w *KernelFactory) Free() error {
|
||||||
if w.ebpfProxy == nil {
|
if w.loopbackProxy == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return w.ebpfProxy.Free()
|
return w.loopbackProxy.Free()
|
||||||
}
|
}
|
||||||
|
|
||||||
func isEBPFDisabled() bool {
|
func isKernelProxyDisabled() bool {
|
||||||
val := os.Getenv(envDisableEBPFWGProxy)
|
env := envDisableKernelWGProxy
|
||||||
|
val := os.Getenv(env)
|
||||||
|
if val == "" {
|
||||||
|
env = envDisableEBPFWGProxy
|
||||||
|
val = os.Getenv(env)
|
||||||
|
}
|
||||||
if val == "" {
|
if val == "" {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
disabled, err := strconv.ParseBool(val)
|
disabled, err := strconv.ParseBool(val)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Warnf("failed to parse %s: %v", envDisableEBPFWGProxy, err)
|
log.Warnf("failed to parse %s: %v", env, err)
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if disabled {
|
||||||
|
log.Infof("kernel WireGuard proxy is disabled via %s", env)
|
||||||
|
}
|
||||||
return disabled
|
return disabled
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -24,11 +24,6 @@ func (w *USPFactory) GetProxy() Proxy {
|
|||||||
return proxyBind.NewProxyBind(w.bind, w.mtu)
|
return proxyBind.NewProxyBind(w.bind, w.mtu)
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetProxyPort returns 0 as userspace WireGuard doesn't use a separate proxy port.
|
|
||||||
func (w *USPFactory) GetProxyPort() uint16 {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *USPFactory) Free() error {
|
func (w *USPFactory) Free() error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,291 @@
|
|||||||
|
//go:build linux && !android
|
||||||
|
|
||||||
|
package loopback
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"sync"
|
||||||
|
"syscall"
|
||||||
|
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
"golang.org/x/net/ipv4"
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
"github.com/netbirdio/netbird/client/iface/bufsize"
|
||||||
|
"github.com/netbirdio/netbird/client/iface/wgproxy/rawsocket"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
loopbackDevice = "lo"
|
||||||
|
|
||||||
|
portRangeStart = 3128
|
||||||
|
portRangeEnd = portRangeStart + 100
|
||||||
|
)
|
||||||
|
|
||||||
|
// Proxy forwards packets between relayed connections and a local kernel
|
||||||
|
// WireGuard instance. Every relayed peer gets its own loopback address as its
|
||||||
|
// WireGuard endpoint, so a single socket serves all of them: the destination
|
||||||
|
// address of an incoming packet identifies the peer.
|
||||||
|
type Proxy struct {
|
||||||
|
localWGListenPort int
|
||||||
|
mtu uint16
|
||||||
|
proxyPort int
|
||||||
|
|
||||||
|
conn *net.UDPConn
|
||||||
|
packetConn *ipv4.PacketConn
|
||||||
|
loIndex int
|
||||||
|
rawConnIPv4 net.PacketConn
|
||||||
|
rawConnIPv6 net.PacketConn
|
||||||
|
|
||||||
|
relayedConnMutex sync.Mutex
|
||||||
|
relayedConnStore map[netip.Addr]net.Conn
|
||||||
|
addrs allocator
|
||||||
|
|
||||||
|
ctx context.Context
|
||||||
|
ctxCancel context.CancelFunc
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewProxy creates a proxy for the WireGuard instance listening on wgPort.
|
||||||
|
func NewProxy(wgPort int, mtu uint16) *Proxy {
|
||||||
|
log.Debugf("instantiate loopback wg proxy")
|
||||||
|
return &Proxy{
|
||||||
|
localWGListenPort: wgPort,
|
||||||
|
mtu: mtu,
|
||||||
|
relayedConnStore: make(map[netip.Addr]net.Conn),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Listen opens the shared socket and starts forwarding WireGuard packets to the
|
||||||
|
// relayed connections.
|
||||||
|
func (p *Proxy) Listen() error {
|
||||||
|
rawConnIPv4, err := rawsocket.PrepareSenderRawSocketIPv4()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("prepare IPv4 raw socket: %w", err)
|
||||||
|
}
|
||||||
|
p.rawConnIPv4 = rawConnIPv4
|
||||||
|
|
||||||
|
p.rawConnIPv6, err = rawsocket.PrepareSenderRawSocketIPv6()
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("failed to prepare IPv6 raw socket, continuing with IPv4 only: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
loopback, err := net.InterfaceByName(loopbackDevice)
|
||||||
|
if err != nil {
|
||||||
|
if freeErr := p.Free(); freeErr != nil {
|
||||||
|
log.Errorf("failed to free the wgproxy: %s", freeErr)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("look up %s: %w", loopbackDevice, err)
|
||||||
|
}
|
||||||
|
p.loIndex = loopback.Index
|
||||||
|
|
||||||
|
if err := p.listen(); err != nil {
|
||||||
|
if freeErr := p.Free(); freeErr != nil {
|
||||||
|
log.Errorf("failed to free the wgproxy: %s", freeErr)
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
p.ctx, p.ctxCancel = context.WithCancel(context.Background())
|
||||||
|
|
||||||
|
go p.proxyToRemote()
|
||||||
|
log.Infof("local wg proxy listening on %s:%d", addrRangePrefix, p.proxyPort)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// listen binds the shared socket on the first free port of the range. The bind
|
||||||
|
// has to be a wildcard one to receive every peer address in the range, so it is
|
||||||
|
// restricted to the loopback device: without that the port would be reachable
|
||||||
|
// on every interface.
|
||||||
|
func (p *Proxy) listen() error {
|
||||||
|
var lastErr error
|
||||||
|
for port := portRangeStart; port <= portRangeEnd; port++ {
|
||||||
|
err := p.listenOn(port)
|
||||||
|
if err == nil {
|
||||||
|
p.proxyPort = port
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
lastErr = err
|
||||||
|
}
|
||||||
|
return fmt.Errorf("bind proxy port in range %d-%d: %w", portRangeStart, portRangeEnd, lastErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Proxy) listenOn(proxyPort int) error {
|
||||||
|
lc := net.ListenConfig{
|
||||||
|
Control: func(_, _ string, c syscall.RawConn) error {
|
||||||
|
var sockErr error
|
||||||
|
if err := c.Control(func(fd uintptr) {
|
||||||
|
if err := unix.SetsockoptString(int(fd), unix.SOL_SOCKET, unix.SO_BINDTODEVICE, loopbackDevice); err != nil {
|
||||||
|
sockErr = fmt.Errorf("bind to %s: %w", loopbackDevice, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}); err != nil {
|
||||||
|
return fmt.Errorf("control socket: %w", err)
|
||||||
|
}
|
||||||
|
return sockErr
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
conn, err := lc.ListenPacket(context.Background(), "udp4", fmt.Sprintf(":%d", proxyPort))
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("listen on :%d: %w", proxyPort, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
udpConn, ok := conn.(*net.UDPConn)
|
||||||
|
if !ok {
|
||||||
|
if closeErr := conn.Close(); closeErr != nil {
|
||||||
|
log.Errorf("failed to close proxy conn: %s", closeErr)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("unexpected conn type %T", conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
packetConn := ipv4.NewPacketConn(udpConn)
|
||||||
|
// the destination address carries the peer identity, the interface index is
|
||||||
|
// checked on receive as a second line of defense behind SO_BINDTODEVICE
|
||||||
|
if err := packetConn.SetControlMessage(ipv4.FlagDst|ipv4.FlagInterface, true); err != nil {
|
||||||
|
if closeErr := udpConn.Close(); closeErr != nil {
|
||||||
|
log.Errorf("failed to close proxy conn: %s", closeErr)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("request destination address: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
p.conn = udpConn
|
||||||
|
p.packetConn = packetConn
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddRelayedConn assigns an endpoint address to the relayed connection and
|
||||||
|
// returns the address WireGuard should send to, along with the key the
|
||||||
|
// connection is stored under.
|
||||||
|
func (p *Proxy) AddRelayedConn(relayedConn net.Conn) (*net.UDPAddr, netip.Addr, error) {
|
||||||
|
addr, err := p.storeRelayedConn(relayedConn)
|
||||||
|
if err != nil {
|
||||||
|
return nil, netip.Addr{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Infof("relayed conn added to wg proxy store: %s, endpoint address: %s", relayedConn.RemoteAddr(), addr)
|
||||||
|
|
||||||
|
return &net.UDPAddr{
|
||||||
|
IP: addr.AsSlice(),
|
||||||
|
Port: p.proxyPort,
|
||||||
|
}, addr, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Free releases the proxy resources. The relayed connections are left open.
|
||||||
|
func (p *Proxy) Free() error {
|
||||||
|
log.Debugf("free up loopback wg proxy")
|
||||||
|
if p.ctx != nil && p.ctx.Err() != nil {
|
||||||
|
//nolint
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.ctxCancel != nil {
|
||||||
|
p.ctxCancel()
|
||||||
|
}
|
||||||
|
|
||||||
|
var result *multierror.Error
|
||||||
|
if p.conn != nil {
|
||||||
|
if err := p.conn.Close(); err != nil {
|
||||||
|
result = multierror.Append(result, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.rawConnIPv4 != nil {
|
||||||
|
if err := p.rawConnIPv4.Close(); err != nil {
|
||||||
|
result = multierror.Append(result, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.rawConnIPv6 != nil {
|
||||||
|
if err := p.rawConnIPv6.Close(); err != nil {
|
||||||
|
result = multierror.Append(result, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nberrors.FormatErrorOrNil(result)
|
||||||
|
}
|
||||||
|
|
||||||
|
// proxyToRemote reads packets from the local WireGuard instance and forwards
|
||||||
|
// them to the relayed connection the destination address belongs to.
|
||||||
|
func (p *Proxy) proxyToRemote() {
|
||||||
|
buf := make([]byte, p.mtu+bufsize.WGBufferOverhead)
|
||||||
|
for p.ctx.Err() == nil {
|
||||||
|
if err := p.readAndForwardPacket(buf); err != nil {
|
||||||
|
if p.ctx.Err() != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
log.Errorf("failed to proxy packet to remote conn: %s", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Proxy) readAndForwardPacket(buf []byte) error {
|
||||||
|
n, cm, _, err := p.packetConn.ReadFrom(buf)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("read UDP packet from WG: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if cm == nil {
|
||||||
|
return fmt.Errorf("no control message on packet")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cm.IfIndex != p.loIndex {
|
||||||
|
log.Tracef("dropping packet received on interface %d instead of %s", cm.IfIndex, loopbackDevice)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
dst, ok := netip.AddrFromSlice(cm.Dst.To4())
|
||||||
|
if !ok || !inRange(dst) {
|
||||||
|
log.Tracef("dropping packet for unexpected destination %s", cm.Dst)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
p.relayedConnMutex.Lock()
|
||||||
|
conn, ok := p.relayedConnStore[dst]
|
||||||
|
p.relayedConnMutex.Unlock()
|
||||||
|
if !ok {
|
||||||
|
if p.ctx.Err() == nil {
|
||||||
|
log.Debugf("relayed conn not found by address because conn already has been closed: %s", dst)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := conn.Write(buf[:n]); err != nil {
|
||||||
|
return fmt.Errorf("forward local WG packet (%s) to remote relayed conn: %w", dst, err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Proxy) storeRelayedConn(relayedConn net.Conn) (netip.Addr, error) {
|
||||||
|
p.relayedConnMutex.Lock()
|
||||||
|
defer p.relayedConnMutex.Unlock()
|
||||||
|
|
||||||
|
addr, err := p.addrs.next(func(a netip.Addr) bool {
|
||||||
|
_, ok := p.relayedConnStore[a]
|
||||||
|
return ok
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return netip.Addr{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
p.relayedConnStore[addr] = relayedConn
|
||||||
|
return addr, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// removeRelayedConn releases an endpoint address. It only removes the entry
|
||||||
|
// while it still belongs to relayedConn, so a late release cannot take an
|
||||||
|
// address away from the peer it was handed to next.
|
||||||
|
func (p *Proxy) removeRelayedConn(addr netip.Addr, relayedConn net.Conn) {
|
||||||
|
p.relayedConnMutex.Lock()
|
||||||
|
defer p.relayedConnMutex.Unlock()
|
||||||
|
|
||||||
|
if stored, ok := p.relayedConnStore[addr]; !ok || stored != relayedConn {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugf("remove relayed conn from store by address: %s", addr)
|
||||||
|
delete(p.relayedConnStore, addr)
|
||||||
|
}
|
||||||
@@ -0,0 +1,196 @@
|
|||||||
|
//go:build linux && !android && privileged
|
||||||
|
|
||||||
|
package loopback
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net"
|
||||||
|
"strconv"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
const testWGPort = 51862
|
||||||
|
|
||||||
|
// relayEnd stands in for a relayed connection: the proxy writes what it read
|
||||||
|
// from WireGuard into it, and the test reads it back out here.
|
||||||
|
func relayEnd(t *testing.T) (proxySide net.Conn, testSide *net.UDPConn) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
testSide, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("relay listener: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
if err := testSide.Close(); err != nil {
|
||||||
|
t.Logf("close relay listener: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
proxySide, err = net.Dial("udp", testSide.LocalAddr().String())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("relay conn: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
if err := proxySide.Close(); err != nil {
|
||||||
|
t.Logf("close relay conn: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
return proxySide, testSide
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestProxyDemuxesByDestinationAddress is the core of the design: one socket
|
||||||
|
// serves every peer, and the destination address decides which relayed
|
||||||
|
// connection a WireGuard packet belongs to.
|
||||||
|
func TestProxyDemuxesByDestinationAddress(t *testing.T) {
|
||||||
|
proxy := NewProxy(testWGPort, 1280)
|
||||||
|
if err := proxy.Listen(); err != nil {
|
||||||
|
t.Fatalf("listen: %v", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if err := proxy.Free(); err != nil {
|
||||||
|
t.Errorf("free proxy: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
const peers = 3
|
||||||
|
endpoints := make([]*net.UDPAddr, 0, peers)
|
||||||
|
readers := make([]*net.UDPConn, 0, peers)
|
||||||
|
for i := 0; i < peers; i++ {
|
||||||
|
proxySide, testSide := relayEnd(t)
|
||||||
|
endpoint, _, err := proxy.AddRelayedConn(proxySide)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("add relayed conn %d: %v", i, err)
|
||||||
|
}
|
||||||
|
if endpoint.Port != proxy.proxyPort {
|
||||||
|
t.Errorf("peer %d endpoint port = %d, want the shared proxy port %d", i, endpoint.Port, proxy.proxyPort)
|
||||||
|
}
|
||||||
|
endpoints = append(endpoints, endpoint)
|
||||||
|
readers = append(readers, testSide)
|
||||||
|
}
|
||||||
|
|
||||||
|
// every peer must have its own address, otherwise they are indistinguishable
|
||||||
|
seen := make(map[string]bool, peers)
|
||||||
|
for i, endpoint := range endpoints {
|
||||||
|
if seen[endpoint.IP.String()] {
|
||||||
|
t.Fatalf("peer %d reuses endpoint address %s", i, endpoint.IP)
|
||||||
|
}
|
||||||
|
seen[endpoint.IP.String()] = true
|
||||||
|
}
|
||||||
|
|
||||||
|
wgSock, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("wg socket: %v", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if err := wgSock.Close(); err != nil {
|
||||||
|
t.Logf("close wg socket: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
for i, endpoint := range endpoints {
|
||||||
|
payload := []byte{byte(i), 'p', 'k', 't'}
|
||||||
|
if _, err := wgSock.WriteTo(payload, endpoint); err != nil {
|
||||||
|
t.Fatalf("write to peer %d endpoint %s: %v", i, endpoint, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
buf := make([]byte, 1500)
|
||||||
|
if err := readers[i].SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil {
|
||||||
|
t.Fatalf("set read deadline: %v", err)
|
||||||
|
}
|
||||||
|
n, _, err := readers[i].ReadFrom(buf)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("peer %d did not receive its packet: %v", i, err)
|
||||||
|
}
|
||||||
|
if string(buf[:n]) != string(payload) {
|
||||||
|
t.Errorf("peer %d got %q, want %q", i, buf[:n], payload)
|
||||||
|
}
|
||||||
|
|
||||||
|
// no other peer may see it
|
||||||
|
for j, other := range readers {
|
||||||
|
if j == i {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := other.SetReadDeadline(time.Now().Add(100 * time.Millisecond)); err != nil {
|
||||||
|
t.Fatalf("set read deadline: %v", err)
|
||||||
|
}
|
||||||
|
if _, _, err := other.ReadFrom(buf); err == nil {
|
||||||
|
t.Errorf("packet for peer %d also delivered to peer %d", i, j)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestProxyDropsPacketsOutsideTheRange guards the wildcard bind: anything that
|
||||||
|
// is not addressed to a handed-out endpoint must not reach a relayed peer.
|
||||||
|
func TestProxyDropsPacketsOutsideTheRange(t *testing.T) {
|
||||||
|
proxy := NewProxy(testWGPort+1, 1280)
|
||||||
|
if err := proxy.Listen(); err != nil {
|
||||||
|
t.Fatalf("listen: %v", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if err := proxy.Free(); err != nil {
|
||||||
|
t.Errorf("free proxy: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
proxySide, testSide := relayEnd(t)
|
||||||
|
if _, _, err := proxy.AddRelayedConn(proxySide); err != nil {
|
||||||
|
t.Fatalf("add relayed conn: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
sender, err := net.Dial("udp", net.JoinHostPort("127.0.0.1", strconv.Itoa(proxy.proxyPort)))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("sender: %v", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if err := sender.Close(); err != nil {
|
||||||
|
t.Logf("close sender: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
if _, err := sender.Write([]byte("stray")); err != nil {
|
||||||
|
t.Fatalf("write stray packet: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
buf := make([]byte, 1500)
|
||||||
|
if err := testSide.SetReadDeadline(time.Now().Add(500 * time.Millisecond)); err != nil {
|
||||||
|
t.Fatalf("set read deadline: %v", err)
|
||||||
|
}
|
||||||
|
if _, _, err := testSide.ReadFrom(buf); err == nil {
|
||||||
|
t.Error("packet addressed to 127.0.0.1 was forwarded to a relayed peer")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A wrapper that is closed before it starts forwarding still has to give its
|
||||||
|
// endpoint address back, otherwise the range leaks an address per attempt.
|
||||||
|
func TestClosingBeforeWorkReleasesTheAddress(t *testing.T) {
|
||||||
|
proxy := NewProxy(testWGPort+2, 1280)
|
||||||
|
if err := proxy.Listen(); err != nil {
|
||||||
|
t.Fatalf("listen: %v", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if err := proxy.Free(); err != nil {
|
||||||
|
t.Errorf("free proxy: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
proxySide, _ := relayEnd(t)
|
||||||
|
wrapper := NewProxyWrapper(proxy)
|
||||||
|
if err := wrapper.AddRelayedConn(context.Background(), nil, proxySide); err != nil {
|
||||||
|
t.Fatalf("add relayed conn: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := len(proxy.relayedConnStore); got != 1 {
|
||||||
|
t.Fatalf("store holds %d entries after adding one conn, want 1", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := wrapper.CloseConn(); err != nil {
|
||||||
|
t.Fatalf("close conn: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := len(proxy.relayedConnStore); got != 0 {
|
||||||
|
t.Errorf("store holds %d entries after close, want 0", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,6 +1,6 @@
|
|||||||
//go:build linux && !android
|
//go:build linux && !android
|
||||||
|
|
||||||
package ebpf
|
package loopback
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
|
"net/netip"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/google/gopacket"
|
"github.com/google/gopacket"
|
||||||
@@ -95,13 +96,14 @@ func NewPacketHeaders(localWGListenPort int, endpoint *net.UDPAddr) (*PacketHead
|
|||||||
|
|
||||||
// ProxyWrapper help to keep the remoteConn instance for net.Conn.Close function call
|
// ProxyWrapper help to keep the remoteConn instance for net.Conn.Close function call
|
||||||
type ProxyWrapper struct {
|
type ProxyWrapper struct {
|
||||||
wgeBPFProxy *WGEBPFProxy
|
proxy *Proxy
|
||||||
|
|
||||||
remoteConn net.Conn
|
remoteConn net.Conn
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
|
|
||||||
wgRelayedEndpointAddr *net.UDPAddr
|
wgRelayedEndpointAddr *net.UDPAddr
|
||||||
|
peerAddr netip.Addr
|
||||||
headers *PacketHeaders
|
headers *PacketHeaders
|
||||||
headerCurrentUsed *PacketHeaders
|
headerCurrentUsed *PacketHeaders
|
||||||
rawConn net.PacketConn
|
rawConn net.PacketConn
|
||||||
@@ -113,36 +115,44 @@ type ProxyWrapper struct {
|
|||||||
closeListener *listener.CloseListener
|
closeListener *listener.CloseListener
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewProxyWrapper(proxy *WGEBPFProxy) *ProxyWrapper {
|
func NewProxyWrapper(proxy *Proxy) *ProxyWrapper {
|
||||||
return &ProxyWrapper{
|
return &ProxyWrapper{
|
||||||
wgeBPFProxy: proxy,
|
proxy: proxy,
|
||||||
pausedCond: sync.NewCond(&sync.Mutex{}),
|
pausedCond: sync.NewCond(&sync.Mutex{}),
|
||||||
closeListener: listener.NewCloseListener(),
|
closeListener: listener.NewCloseListener(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *ProxyWrapper) AddRelayedConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error {
|
func (p *ProxyWrapper) AddRelayedConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error {
|
||||||
addr, err := p.wgeBPFProxy.AddRelayedConn(remoteConn)
|
addr, peerAddr, err := p.proxy.AddRelayedConn(remoteConn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("add relayed conn: %w", err)
|
return fmt.Errorf("add relayed conn: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
headers, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, addr)
|
// the endpoint address is otherwise only released by the forwarding
|
||||||
|
// goroutine, which never starts when the setup below fails
|
||||||
|
release := func() { p.proxy.removeRelayedConn(peerAddr, remoteConn) }
|
||||||
|
|
||||||
|
headers, err := NewPacketHeaders(p.proxy.localWGListenPort, addr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
release()
|
||||||
return fmt.Errorf("create packet sender: %w", err)
|
return fmt.Errorf("create packet sender: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check if required raw connection is available
|
// Check if required raw connection is available
|
||||||
if !headers.isIPv4 && p.wgeBPFProxy.rawConnIPv6 == nil {
|
if !headers.isIPv4 && p.proxy.rawConnIPv6 == nil {
|
||||||
|
release()
|
||||||
return errIPv6ConnNotAvailable
|
return errIPv6ConnNotAvailable
|
||||||
}
|
}
|
||||||
if headers.isIPv4 && p.wgeBPFProxy.rawConnIPv4 == nil {
|
if headers.isIPv4 && p.proxy.rawConnIPv4 == nil {
|
||||||
|
release()
|
||||||
return errIPv4ConnNotAvailable
|
return errIPv4ConnNotAvailable
|
||||||
}
|
}
|
||||||
|
|
||||||
p.remoteConn = remoteConn
|
p.remoteConn = remoteConn
|
||||||
p.ctx, p.cancel = context.WithCancel(ctx)
|
p.ctx, p.cancel = context.WithCancel(ctx)
|
||||||
p.wgRelayedEndpointAddr = addr
|
p.wgRelayedEndpointAddr = addr
|
||||||
|
p.peerAddr = peerAddr
|
||||||
p.headers = headers
|
p.headers = headers
|
||||||
p.rawConn = p.selectRawConn(headers)
|
p.rawConn = p.selectRawConn(headers)
|
||||||
return nil
|
return nil
|
||||||
@@ -193,18 +203,18 @@ func (p *ProxyWrapper) RedirectAs(endpoint *net.UDPAddr) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
header, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, endpoint)
|
header, err := NewPacketHeaders(p.proxy.localWGListenPort, endpoint)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Errorf("failed to create packet headers: %s", err)
|
log.Errorf("failed to create packet headers: %s", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check if required raw connection is available
|
// Check if required raw connection is available
|
||||||
if !header.isIPv4 && p.wgeBPFProxy.rawConnIPv6 == nil {
|
if !header.isIPv4 && p.proxy.rawConnIPv6 == nil {
|
||||||
log.Error(errIPv6ConnNotAvailable)
|
log.Error(errIPv6ConnNotAvailable)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if header.isIPv4 && p.wgeBPFProxy.rawConnIPv4 == nil {
|
if header.isIPv4 && p.proxy.rawConnIPv4 == nil {
|
||||||
log.Error(errIPv4ConnNotAvailable)
|
log.Error(errIPv4ConnNotAvailable)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -240,6 +250,10 @@ func (p *ProxyWrapper) CloseConn() error {
|
|||||||
|
|
||||||
p.closeListener.SetCloseListener(nil)
|
p.closeListener.SetCloseListener(nil)
|
||||||
|
|
||||||
|
// releases the endpoint address for a wrapper that was never started, and
|
||||||
|
// is a no-op once the forwarding goroutine has released it
|
||||||
|
p.proxy.removeRelayedConn(p.peerAddr, p.remoteConn)
|
||||||
|
|
||||||
p.pausedCond.L.Lock()
|
p.pausedCond.L.Lock()
|
||||||
p.paused = false
|
p.paused = false
|
||||||
p.pausedCond.Signal()
|
p.pausedCond.Signal()
|
||||||
@@ -252,9 +266,9 @@ func (p *ProxyWrapper) CloseConn() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (p *ProxyWrapper) proxyToLocal(ctx context.Context) {
|
func (p *ProxyWrapper) proxyToLocal(ctx context.Context) {
|
||||||
defer p.wgeBPFProxy.removeRelayedConn(uint16(p.wgRelayedEndpointAddr.Port))
|
defer p.proxy.removeRelayedConn(p.peerAddr, p.remoteConn)
|
||||||
|
|
||||||
buf := make([]byte, p.wgeBPFProxy.mtu+bufsize.WGBufferOverhead)
|
buf := make([]byte, p.proxy.mtu+bufsize.WGBufferOverhead)
|
||||||
for {
|
for {
|
||||||
n, err := p.readFromRemote(ctx, buf)
|
n, err := p.readFromRemote(ctx, buf)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -286,7 +300,7 @@ func (p *ProxyWrapper) readFromRemote(ctx context.Context, buf []byte) (int, err
|
|||||||
}
|
}
|
||||||
p.closeListener.Notify()
|
p.closeListener.Notify()
|
||||||
if !errors.Is(err, io.EOF) {
|
if !errors.Is(err, io.EOF) {
|
||||||
log.Errorf("failed to read from relayed conn (endpoint: :%d): %s", p.wgRelayedEndpointAddr.Port, err)
|
log.Errorf("failed to read from relayed conn (endpoint: %s): %s", p.wgRelayedEndpointAddr, err)
|
||||||
}
|
}
|
||||||
return 0, err
|
return 0, err
|
||||||
}
|
}
|
||||||
@@ -314,7 +328,7 @@ func (p *ProxyWrapper) sendPkg(data []byte, header *PacketHeaders) error {
|
|||||||
|
|
||||||
func (p *ProxyWrapper) selectRawConn(header *PacketHeaders) net.PacketConn {
|
func (p *ProxyWrapper) selectRawConn(header *PacketHeaders) net.PacketConn {
|
||||||
if header.isIPv4 {
|
if header.isIPv4 {
|
||||||
return p.wgeBPFProxy.rawConnIPv4
|
return p.proxy.rawConnIPv4
|
||||||
}
|
}
|
||||||
return p.wgeBPFProxy.rawConnIPv6
|
return p.proxy.rawConnIPv6
|
||||||
}
|
}
|
||||||
@@ -9,25 +9,25 @@ import (
|
|||||||
"github.com/netbirdio/netbird/client/iface/bind"
|
"github.com/netbirdio/netbird/client/iface/bind"
|
||||||
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
||||||
bindproxy "github.com/netbirdio/netbird/client/iface/wgproxy/bind"
|
bindproxy "github.com/netbirdio/netbird/client/iface/wgproxy/bind"
|
||||||
"github.com/netbirdio/netbird/client/iface/wgproxy/ebpf"
|
"github.com/netbirdio/netbird/client/iface/wgproxy/loopback"
|
||||||
"github.com/netbirdio/netbird/client/iface/wgproxy/udp"
|
"github.com/netbirdio/netbird/client/iface/wgproxy/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
func seedProxies() ([]proxyInstance, error) {
|
func seedProxies() ([]proxyInstance, error) {
|
||||||
pl := make([]proxyInstance, 0)
|
pl := make([]proxyInstance, 0)
|
||||||
|
|
||||||
ebpfProxy := ebpf.NewWGEBPFProxy(51831, 1280)
|
loopbackProxy := loopback.NewProxy(51831, 1280)
|
||||||
if err := ebpfProxy.Listen(); err != nil {
|
if err := loopbackProxy.Listen(); err != nil {
|
||||||
return nil, fmt.Errorf("failed to initialize ebpf proxy: %s", err)
|
return nil, fmt.Errorf("failed to initialize loopback proxy: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
pEbpf := proxyInstance{
|
pLoopback := proxyInstance{
|
||||||
name: "ebpf kernel proxy",
|
name: "loopback kernel proxy",
|
||||||
proxy: ebpf.NewProxyWrapper(ebpfProxy),
|
proxy: loopback.NewProxyWrapper(loopbackProxy),
|
||||||
wgPort: 51831,
|
wgPort: 51831,
|
||||||
closeFn: ebpfProxy.Free,
|
closeFn: loopbackProxy.Free,
|
||||||
}
|
}
|
||||||
pl = append(pl, pEbpf)
|
pl = append(pl, pLoopback)
|
||||||
|
|
||||||
pUDP := proxyInstance{
|
pUDP := proxyInstance{
|
||||||
name: "udp kernel proxy",
|
name: "udp kernel proxy",
|
||||||
@@ -42,18 +42,18 @@ func seedProxies() ([]proxyInstance, error) {
|
|||||||
func seedProxyForProxyCloseByRemoteConn() ([]proxyInstance, error) {
|
func seedProxyForProxyCloseByRemoteConn() ([]proxyInstance, error) {
|
||||||
pl := make([]proxyInstance, 0)
|
pl := make([]proxyInstance, 0)
|
||||||
|
|
||||||
ebpfProxy := ebpf.NewWGEBPFProxy(51831, 1280)
|
loopbackProxy := loopback.NewProxy(51831, 1280)
|
||||||
if err := ebpfProxy.Listen(); err != nil {
|
if err := loopbackProxy.Listen(); err != nil {
|
||||||
return nil, fmt.Errorf("failed to initialize ebpf proxy: %s", err)
|
return nil, fmt.Errorf("failed to initialize loopback proxy: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
pEbpf := proxyInstance{
|
pLoopback := proxyInstance{
|
||||||
name: "ebpf kernel proxy",
|
name: "loopback kernel proxy",
|
||||||
proxy: ebpf.NewProxyWrapper(ebpfProxy),
|
proxy: loopback.NewProxyWrapper(loopbackProxy),
|
||||||
wgPort: 51831,
|
wgPort: 51831,
|
||||||
closeFn: ebpfProxy.Free,
|
closeFn: loopbackProxy.Free,
|
||||||
}
|
}
|
||||||
pl = append(pl, pEbpf)
|
pl = append(pl, pLoopback)
|
||||||
|
|
||||||
pUDP := proxyInstance{
|
pUDP := proxyInstance{
|
||||||
name: "udp kernel proxy",
|
name: "udp kernel proxy",
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/iface/wgproxy/ebpf"
|
"github.com/netbirdio/netbird/client/iface/wgproxy/loopback"
|
||||||
"github.com/netbirdio/netbird/client/iface/wgproxy/udp"
|
"github.com/netbirdio/netbird/client/iface/wgproxy/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -198,20 +198,20 @@ func testRedirectAs(t *testing.T, proxy Proxy, wgPort int, nbAddr, p2pEndpoint *
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestRedirectAs_eBPF_IPv4 tests RedirectAs with eBPF proxy using IPv4 addresses
|
// TestRedirectAs_Loopback_IPv4 tests RedirectAs with the loopback proxy using IPv4 addresses
|
||||||
func TestRedirectAs_eBPF_IPv4(t *testing.T) {
|
func TestRedirectAs_Loopback_IPv4(t *testing.T) {
|
||||||
wgPort := 51850
|
wgPort := 51850
|
||||||
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280)
|
loopbackProxy := loopback.NewProxy(wgPort, 1280)
|
||||||
if err := ebpfProxy.Listen(); err != nil {
|
if err := loopbackProxy.Listen(); err != nil {
|
||||||
t.Fatalf("failed to initialize ebpf proxy: %v", err)
|
t.Fatalf("failed to initialize loopback proxy: %v", err)
|
||||||
}
|
}
|
||||||
defer func() {
|
defer func() {
|
||||||
if err := ebpfProxy.Free(); err != nil {
|
if err := loopbackProxy.Free(); err != nil {
|
||||||
t.Errorf("failed to free ebpf proxy: %v", err)
|
t.Errorf("failed to free loopback proxy: %v", err)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
proxy := ebpf.NewProxyWrapper(ebpfProxy)
|
proxy := loopback.NewProxyWrapper(loopbackProxy)
|
||||||
|
|
||||||
// NetBird UDP address of the remote peer
|
// NetBird UDP address of the remote peer
|
||||||
nbAddr := &net.UDPAddr{
|
nbAddr := &net.UDPAddr{
|
||||||
@@ -227,20 +227,20 @@ func TestRedirectAs_eBPF_IPv4(t *testing.T) {
|
|||||||
testRedirectAs(t, proxy, wgPort, nbAddr, p2pEndpoint)
|
testRedirectAs(t, proxy, wgPort, nbAddr, p2pEndpoint)
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestRedirectAs_eBPF_IPv6 tests RedirectAs with eBPF proxy using IPv6 addresses
|
// TestRedirectAs_Loopback_IPv6 tests RedirectAs with the loopback proxy using IPv6 addresses
|
||||||
func TestRedirectAs_eBPF_IPv6(t *testing.T) {
|
func TestRedirectAs_Loopback_IPv6(t *testing.T) {
|
||||||
wgPort := 51851
|
wgPort := 51851
|
||||||
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280)
|
loopbackProxy := loopback.NewProxy(wgPort, 1280)
|
||||||
if err := ebpfProxy.Listen(); err != nil {
|
if err := loopbackProxy.Listen(); err != nil {
|
||||||
t.Fatalf("failed to initialize ebpf proxy: %v", err)
|
t.Fatalf("failed to initialize loopback proxy: %v", err)
|
||||||
}
|
}
|
||||||
defer func() {
|
defer func() {
|
||||||
if err := ebpfProxy.Free(); err != nil {
|
if err := loopbackProxy.Free(); err != nil {
|
||||||
t.Errorf("failed to free ebpf proxy: %v", err)
|
t.Errorf("failed to free loopback proxy: %v", err)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
proxy := ebpf.NewProxyWrapper(ebpfProxy)
|
proxy := loopback.NewProxyWrapper(loopbackProxy)
|
||||||
|
|
||||||
// NetBird UDP address of the remote peer
|
// NetBird UDP address of the remote peer
|
||||||
nbAddr := &net.UDPAddr{
|
nbAddr := &net.UDPAddr{
|
||||||
@@ -259,17 +259,17 @@ func TestRedirectAs_eBPF_IPv6(t *testing.T) {
|
|||||||
// TestRedirectAs_Multiple_Switches tests switching between multiple endpoints
|
// TestRedirectAs_Multiple_Switches tests switching between multiple endpoints
|
||||||
func TestRedirectAs_Multiple_Switches(t *testing.T) {
|
func TestRedirectAs_Multiple_Switches(t *testing.T) {
|
||||||
wgPort := 51856
|
wgPort := 51856
|
||||||
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280)
|
loopbackProxy := loopback.NewProxy(wgPort, 1280)
|
||||||
if err := ebpfProxy.Listen(); err != nil {
|
if err := loopbackProxy.Listen(); err != nil {
|
||||||
t.Fatalf("failed to initialize ebpf proxy: %v", err)
|
t.Fatalf("failed to initialize loopback proxy: %v", err)
|
||||||
}
|
}
|
||||||
defer func() {
|
defer func() {
|
||||||
if err := ebpfProxy.Free(); err != nil {
|
if err := loopbackProxy.Free(); err != nil {
|
||||||
t.Errorf("failed to free ebpf proxy: %v", err)
|
t.Errorf("failed to free loopback proxy: %v", err)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
proxy := ebpf.NewProxyWrapper(ebpfProxy)
|
proxy := loopback.NewProxyWrapper(loopbackProxy)
|
||||||
|
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
|
|||||||
@@ -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"
|
||||||
|
}
|
||||||
@@ -103,7 +103,7 @@ func (a *Auth) IsSSOSupported(ctx context.Context) (bool, error) {
|
|||||||
|
|
||||||
err := a.withRetry(ctx, func(client *mgm.GrpcClient) error {
|
err := a.withRetry(ctx, func(client *mgm.GrpcClient) error {
|
||||||
// Try PKCE flow first
|
// Try PKCE flow first
|
||||||
_, err := a.getPKCEFlow(client)
|
_, err := a.getPKCEFlow(client, false)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
supportsSSO = true
|
supportsSSO = true
|
||||||
return nil
|
return nil
|
||||||
@@ -136,9 +136,13 @@ func (a *Auth) IsSSOSupported(ctx context.Context) (bool, error) {
|
|||||||
return supportsSSO, err
|
return supportsSSO, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetOAuthFlow returns an OAuth flow (PKCE or Device) using the existing management connection
|
// GetOAuthFlow returns an OAuth flow (PKCE or Device) using the existing management connection.
|
||||||
// This avoids creating a new connection to the management server
|
// This avoids creating a new connection to the management server.
|
||||||
func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool, hint string) (OAuthFlow, error) {
|
//
|
||||||
|
// 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
|
var flow OAuthFlow
|
||||||
|
|
||||||
err := a.withRetry(ctx, func(client *mgm.GrpcClient) error {
|
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
|
// Try PKCE flow first
|
||||||
pkceFlow, err := a.getPKCEFlow(client)
|
pkceFlow, err := a.getPKCEFlow(client, sessionExtend)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// If PKCE not supported, try Device flow
|
// If PKCE not supported, try Device flow
|
||||||
if s, ok := status.FromError(err); ok && (s.Code() == codes.NotFound || s.Code() == codes.Unimplemented) {
|
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
|
// getPKCEFlow retrieves PKCE authorization flow configuration and creates a flow instance
|
||||||
func (a *Auth) getPKCEFlow(client *mgm.GrpcClient) (*PKCEAuthorizationFlow, error) {
|
func (a *Auth) getPKCEFlow(client *mgm.GrpcClient, sessionExtend bool) (*PKCEAuthorizationFlow, error) {
|
||||||
protoFlow, err := client.GetPKCEAuthorizationFlow()
|
protoFlow, err := client.GetPKCEAuthorizationFlow(sessionExtend)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if s, ok := status.FromError(err); ok && s.Code() == codes.NotFound {
|
if s, ok := status.FromError(err); ok && s.Code() == codes.NotFound {
|
||||||
log.Warnf("server couldn't find pkce flow, contact admin: %v", err)
|
log.Warnf("server couldn't find pkce flow, contact admin: %v", err)
|
||||||
|
|||||||
@@ -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
|
// callers store to send back as the login_hint. Without it a client
|
||||||
// driven through the device flow — Android TV and tvOS — never binds
|
// driven through the device flow — Android TV and tvOS — never binds
|
||||||
// an account to its profile and every later login goes out blind.
|
// 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)
|
log.Warnf("failed to parse email from ID token: %v", err)
|
||||||
} else {
|
} else {
|
||||||
tokenInfo.Email = email
|
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))
|
log.Infof("device flow: user authorization confirmed after %d polls in %s", polls, time.Since(start).Round(time.Second))
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user