mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-08 22:49:10 +02:00
Merge branch 'main' into fix/pkce-flow-session-extend
# Conflicts: # client/ios/NetBirdSDK/login.go # client/server/server.go # shared/management/proto/management.pb.go
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
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
FROM golang:1.25-bookworm
|
FROM golang:1.26.7-bookworm
|
||||||
|
|
||||||
RUN apt-get update && export DEBIAN_FRONTEND=noninteractive \
|
RUN apt-get update && export DEBIAN_FRONTEND=noninteractive \
|
||||||
&& apt-get -y install --no-install-recommends\
|
&& apt-get -y install --no-install-recommends\
|
||||||
|
|||||||
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::"
|
||||||
@@ -27,7 +27,22 @@ jobs:
|
|||||||
push: false
|
push: false
|
||||||
archive: false
|
archive: false
|
||||||
pr_comment: false
|
pr_comment: false
|
||||||
build: false
|
|
||||||
lint: false
|
lint: false
|
||||||
format: false
|
format: false
|
||||||
breaking: true
|
# A push that creates a branch carries no `before` commit, so the
|
||||||
|
# action's default baseline is the all-zero SHA and `buf breaking`
|
||||||
|
# dies cloning it. Skipping costs nothing: every commit on a freshly
|
||||||
|
# cut release branch should have already passed this check on main.
|
||||||
|
breaking: ${{ !github.event.created }}
|
||||||
|
# The alternative is to compare against the default branch instead of
|
||||||
|
# skipping. Not used: buf clones the baseline when the job runs, so a
|
||||||
|
# main that has moved on since the branch was cut reads as protos
|
||||||
|
# deleted on the release branch. Resolving to an empty string on every
|
||||||
|
# other event is what keeps the action's own default in place, which
|
||||||
|
# stacked pull requests need.
|
||||||
|
# breaking_against: >-
|
||||||
|
# ${{ github.event.created
|
||||||
|
# && format('{0}#format=git,branch={1}',
|
||||||
|
# github.event.repository.clone_url,
|
||||||
|
# github.event.repository.default_branch)
|
||||||
|
# || '' }}
|
||||||
|
|||||||
@@ -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 }}
|
||||||
@@ -233,7 +266,7 @@ jobs:
|
|||||||
-e GOCACHE=${CONTAINER_GOCACHE} \
|
-e GOCACHE=${CONTAINER_GOCACHE} \
|
||||||
-e GOMODCACHE=${CONTAINER_GOMODCACHE} \
|
-e GOMODCACHE=${CONTAINER_GOMODCACHE} \
|
||||||
-e CONTAINER=${CONTAINER} \
|
-e CONTAINER=${CONTAINER} \
|
||||||
golang:1.25-alpine \
|
golang:1.26.7-alpine \
|
||||||
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; \
|
||||||
@@ -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$' }
|
||||||
$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
|
||||||
|
|||||||
@@ -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/...
|
||||||
@@ -0,0 +1,199 @@
|
|||||||
|
name: Red Hat Certification
|
||||||
|
|
||||||
|
# Certify published UBI images in the Red Hat Ecosystem Catalog. Called by
|
||||||
|
# release.yml on stable tags, or run by hand to (re)certify any released
|
||||||
|
# version. preflight submits every architecture of an image's manifest list
|
||||||
|
# to Pyxis; auto-publish on the component makes it public once certified.
|
||||||
|
#
|
||||||
|
# Each component's Partner Connect ID comes from the REDHAT_CERT_ID_<NAME>
|
||||||
|
# repository variable, e.g. REDHAT_CERT_ID_CLIENT_ROOTLESS. The run fails
|
||||||
|
# before certifying anything if a selected component's variable is not set.
|
||||||
|
|
||||||
|
on:
|
||||||
|
workflow_call:
|
||||||
|
inputs:
|
||||||
|
component:
|
||||||
|
type: string
|
||||||
|
required: true
|
||||||
|
version:
|
||||||
|
type: string
|
||||||
|
required: true
|
||||||
|
secrets:
|
||||||
|
PYXIS_API_TOKEN:
|
||||||
|
required: true
|
||||||
|
workflow_dispatch:
|
||||||
|
inputs:
|
||||||
|
component:
|
||||||
|
description: "Component to certify"
|
||||||
|
type: choice
|
||||||
|
required: true
|
||||||
|
default: all
|
||||||
|
options:
|
||||||
|
- all
|
||||||
|
- client-rootless
|
||||||
|
- reverse-proxy
|
||||||
|
version:
|
||||||
|
description: "Released version, e.g. v0.80.0"
|
||||||
|
type: string
|
||||||
|
required: true
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
resolve:
|
||||||
|
name: Resolve components
|
||||||
|
runs-on: ubuntu-24.04
|
||||||
|
outputs:
|
||||||
|
version: ${{ steps.resolve.outputs.version }}
|
||||||
|
matrix: ${{ steps.resolve.outputs.matrix }}
|
||||||
|
steps:
|
||||||
|
- name: Resolve components and images
|
||||||
|
id: resolve
|
||||||
|
env:
|
||||||
|
COMPONENT: ${{ inputs.component }}
|
||||||
|
INPUT_VERSION: ${{ inputs.version }}
|
||||||
|
REPO_VARS: ${{ toJSON(vars) }}
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
version="${INPUT_VERSION#v}"
|
||||||
|
if [[ ! "$version" =~ ^[0-9]+\.[0-9]+\.[0-9]+$ ]]; then
|
||||||
|
echo "::error::Only stable x.y.z versions are certified, got '${INPUT_VERSION}'"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
# name, image repository, tag suffix (must match .goreleaser.yaml).
|
||||||
|
# Keep the names in sync with the workflow_dispatch options above.
|
||||||
|
components=(
|
||||||
|
"client-rootless ghcr.io/netbirdio/netbird -rootless-ubi"
|
||||||
|
"reverse-proxy ghcr.io/netbirdio/reverse-proxy -ubi"
|
||||||
|
)
|
||||||
|
matrix="[]"
|
||||||
|
missing=()
|
||||||
|
for c in "${components[@]}"; do
|
||||||
|
read -r name repo suffix <<< "$c"
|
||||||
|
[[ "$COMPONENT" == all || "$COMPONENT" == "$name" ]] || continue
|
||||||
|
var="REDHAT_CERT_ID_${name^^}"; var="${var//-/_}"
|
||||||
|
id="$(jq -r --arg v "$var" '.[$v] // empty' <<< "$REPO_VARS")"
|
||||||
|
if [[ -z "$id" ]]; then
|
||||||
|
missing+=("$var")
|
||||||
|
continue
|
||||||
|
fi
|
||||||
|
matrix="$(jq -c --arg n "$name" --arg t "${version}${suffix}" --arg r "${repo}:${version}${suffix}" --arg i "$id" \
|
||||||
|
'. + [{component: $n, tag: $t, ref: $r, component_id: $i}]' <<< "$matrix")"
|
||||||
|
done
|
||||||
|
if (( ${#missing[@]} )); then
|
||||||
|
echo "::error::Set these repository variables to the Partner Connect component IDs: ${missing[*]}"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
if [[ "$matrix" == "[]" ]]; then
|
||||||
|
echo "::error::No component to certify for '${COMPONENT}'"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
echo "Components to certify: ${matrix}"
|
||||||
|
echo "version=${version}" >> "$GITHUB_OUTPUT"
|
||||||
|
echo "matrix=${matrix}" >> "$GITHUB_OUTPUT"
|
||||||
|
|
||||||
|
certify:
|
||||||
|
name: "Certify ${{ matrix.component }} UBI image"
|
||||||
|
needs: resolve
|
||||||
|
runs-on: ubuntu-24.04
|
||||||
|
strategy:
|
||||||
|
fail-fast: false
|
||||||
|
matrix:
|
||||||
|
include: ${{ fromJSON(needs.resolve.outputs.matrix) }}
|
||||||
|
env:
|
||||||
|
PREFLIGHT_VERSION: "1.21.0"
|
||||||
|
# sha256 of preflight-linux-amd64 from the 1.21.0 GitHub release.
|
||||||
|
# Red Hat publishes no checksum file, so the value is pinned here.
|
||||||
|
PREFLIGHT_SHA256: "5e653135503c72f8702bbe31d7643197d12937c68086879133dd6b9650a9a449"
|
||||||
|
steps:
|
||||||
|
- name: Verify the multi-arch image is on ghcr.io
|
||||||
|
env:
|
||||||
|
IMAGE_REF: ${{ matrix.ref }}
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
docker buildx imagetools inspect "$IMAGE_REF" --raw > manifest.json
|
||||||
|
for arch in amd64 arm64; do
|
||||||
|
if ! jq -e --arg a "$arch" '.manifests[] | select(.platform.architecture == $a)' manifest.json > /dev/null; then
|
||||||
|
echo "::error::${IMAGE_REF} has no ${arch} manifest"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
done
|
||||||
|
echo "Manifest list for ${IMAGE_REF}:"
|
||||||
|
jq -r '.manifests[] | "\(.platform.os)/\(.platform.architecture) \(.digest)"' manifest.json
|
||||||
|
|
||||||
|
- name: Install preflight
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
curl -fsSL --proto '=https' --proto-redir '=https' -o preflight \
|
||||||
|
"https://github.com/redhat-openshift-ecosystem/openshift-preflight/releases/download/${PREFLIGHT_VERSION}/preflight-linux-amd64"
|
||||||
|
echo "${PREFLIGHT_SHA256} preflight" | sha256sum -c -
|
||||||
|
chmod +x preflight
|
||||||
|
./preflight --version
|
||||||
|
|
||||||
|
- name: Run preflight checks and submit to Red Hat
|
||||||
|
env:
|
||||||
|
IMAGE_REF: ${{ matrix.ref }}
|
||||||
|
PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
|
||||||
|
PFLT_CERTIFICATION_COMPONENT_ID: ${{ matrix.component_id }}
|
||||||
|
PFLT_ARTIFACTS: artifacts
|
||||||
|
PFLT_LOGFILE: artifacts/preflight.log
|
||||||
|
PFLT_LOGLEVEL: info
|
||||||
|
PFLT_JUNIT: "true"
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
# No --platform: preflight walks the manifest list and submits every
|
||||||
|
# architecture in one run, grouped under one manifest-list digest.
|
||||||
|
# preflight does not create the PFLT_LOGFILE directory, and --submit
|
||||||
|
# fails if the log file is missing.
|
||||||
|
mkdir -p artifacts
|
||||||
|
./preflight check container "$IMAGE_REF" --submit
|
||||||
|
|
||||||
|
- name: Fail if any check did not pass
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
shopt -s nullglob
|
||||||
|
results=(artifacts/results.json artifacts/*/results.json)
|
||||||
|
if [[ ${#results[@]} -eq 0 ]]; then
|
||||||
|
echo "::error::preflight produced no results.json"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
status=0
|
||||||
|
for f in "${results[@]}"; do
|
||||||
|
arch="$(basename "$(dirname "$f")")"
|
||||||
|
passed="$(jq -r '.passed' "$f")"
|
||||||
|
failed="$(jq -r '[.results.failed[]?.name] | join(", ")' "$f")"
|
||||||
|
echo "${arch}: passed=${passed} ${failed:+failed checks: ${failed}}"
|
||||||
|
[[ "$passed" == "true" ]] || status=1
|
||||||
|
done
|
||||||
|
exit $status
|
||||||
|
|
||||||
|
- name: Upload preflight artifacts
|
||||||
|
if: always()
|
||||||
|
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
|
||||||
|
with:
|
||||||
|
name: redhat-preflight-${{ matrix.component }}-${{ needs.resolve.outputs.version }}
|
||||||
|
path: artifacts/
|
||||||
|
retention-days: 30
|
||||||
|
|
||||||
|
- name: Wait for Pyxis to mark both architectures certified
|
||||||
|
env:
|
||||||
|
TAG: ${{ matrix.tag }}
|
||||||
|
COMPONENT_ID: ${{ matrix.component_id }}
|
||||||
|
PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
# Filter on the tag server-side so older versions are found past the first page.
|
||||||
|
url="https://catalog.redhat.com/api/containers/v1/projects/certification/id/${COMPONENT_ID}/images?filter=repositories.tags.name==${TAG}&page_size=100"
|
||||||
|
for attempt in $(seq 1 20); do
|
||||||
|
certified="$(curl -fsS --proto '=https' --proto-redir '=https' -H "X-API-KEY: ${PFLT_PYXIS_API_TOKEN}" "$url" \
|
||||||
|
| jq -r --arg t "$TAG" '[.data[] | select(.repositories[]?.tags[]?.name == $t) | select(.certified == true) | .architecture] | unique | join(",")')"
|
||||||
|
echo "attempt ${attempt}: certified architectures for ${TAG}: ${certified:-none}"
|
||||||
|
if [[ "$certified" == "amd64,arm64" ]]; then
|
||||||
|
echo "Both architectures certified. Auto-publish is enabled on the component, so the catalog updates on its own."
|
||||||
|
exit 0
|
||||||
|
fi
|
||||||
|
sleep 30
|
||||||
|
done
|
||||||
|
echo "::error::Pyxis has not marked both architectures certified after 10 minutes. Check https://connect.redhat.com/component/view/${COMPONENT_ID}/images"
|
||||||
|
exit 1
|
||||||
@@ -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
|
||||||
|
# proxy/collect-licenses.sh reads the UI's license terms from node_modules.
|
||||||
|
working-directory: proxy/web
|
||||||
|
run: npm ci --ignore-scripts
|
||||||
- name: Set up QEMU
|
- name: Set up QEMU
|
||||||
uses: docker/setup-qemu-action@06116385d9baf250c9f4dcb4858b16962ea869c3 #v4.1.0
|
uses: docker/setup-qemu-action@06116385d9baf250c9f4dcb4858b16962ea869c3 #v4.1.0
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
@@ -215,7 +231,7 @@ jobs:
|
|||||||
echo "GPG_RPM_KEY_FILE=/tmp/gpg-rpm-signing-key.asc" >> $GITHUB_ENV
|
echo "GPG_RPM_KEY_FILE=/tmp/gpg-rpm-signing-key.asc" >> $GITHUB_ENV
|
||||||
|
|
||||||
- name: Install goversioninfo
|
- name: Install goversioninfo
|
||||||
run: go install github.com/josephspurrier/goversioninfo/cmd/goversioninfo@233067e
|
run: go install github.com/josephspurrier/goversioninfo/cmd/goversioninfo@b66839b # v1.7.0
|
||||||
- name: Generate windows syso amd64
|
- name: Generate windows syso amd64
|
||||||
run: goversioninfo -icon client/ui/build/windows/icon.ico -manifest client/manifest.xml -product-name ${{ env.PRODUCT_NAME }} -copyright "${{ env.COPYRIGHT }}" -ver-major ${{ steps.semver_parser.outputs.major }} -ver-minor ${{ steps.semver_parser.outputs.minor }} -ver-patch ${{ steps.semver_parser.outputs.patch }} -ver-build 0 -file-version ${{ steps.semver_parser.outputs.fullversion }}.0 -product-version ${{ steps.semver_parser.outputs.fullversion }}.0 -o client/resources_windows_amd64.syso
|
run: goversioninfo -icon client/ui/build/windows/icon.ico -manifest client/manifest.xml -product-name ${{ env.PRODUCT_NAME }} -copyright "${{ env.COPYRIGHT }}" -ver-major ${{ steps.semver_parser.outputs.major }} -ver-minor ${{ steps.semver_parser.outputs.minor }} -ver-patch ${{ steps.semver_parser.outputs.patch }} -ver-build 0 -file-version ${{ steps.semver_parser.outputs.fullversion }}.0 -product-version ${{ steps.semver_parser.outputs.fullversion }}.0 -o client/resources_windows_amd64.syso
|
||||||
- name: Generate windows syso arm64
|
- name: Generate windows syso arm64
|
||||||
@@ -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
|
||||||
|
|
||||||
@@ -435,7 +480,7 @@ jobs:
|
|||||||
tar -xf llvm-mingw-20250709-ucrt-ubuntu-22.04-x86_64.tar.xz
|
tar -xf llvm-mingw-20250709-ucrt-ubuntu-22.04-x86_64.tar.xz
|
||||||
echo "/tmp/llvm-mingw-20250709-ucrt-ubuntu-22.04-x86_64/bin" >> $GITHUB_PATH
|
echo "/tmp/llvm-mingw-20250709-ucrt-ubuntu-22.04-x86_64/bin" >> $GITHUB_PATH
|
||||||
- name: Install goversioninfo
|
- name: Install goversioninfo
|
||||||
run: go install github.com/josephspurrier/goversioninfo/cmd/goversioninfo@233067e
|
run: go install github.com/josephspurrier/goversioninfo/cmd/goversioninfo@b66839b # v1.7.0
|
||||||
- name: Install wails3 CLI
|
- name: Install wails3 CLI
|
||||||
# Version derived from go.mod so the binding generator always matches
|
# Version derived from go.mod so the binding generator always matches
|
||||||
# the wails runtime the binary links against.
|
# the wails runtime the binary links against.
|
||||||
@@ -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,7 +32,7 @@ jobs:
|
|||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
- name: Set up Node.js
|
- name: Set up Node.js
|
||||||
uses: actions/setup-node@v4
|
uses: actions/setup-node@v7
|
||||||
with:
|
with:
|
||||||
node-version: "22"
|
node-version: "22"
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+156
-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
|
||||||
@@ -223,23 +249,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 +364,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 client/collect-licenses.sh "{{ .ContextDir }}/licenses" amd64 arm64'
|
||||||
|
env:
|
||||||
|
- GOOS=linux
|
||||||
|
- CGO_ENABLED=0
|
||||||
|
labels:
|
||||||
|
"org.opencontainers.image.created": "{{.Date}}"
|
||||||
|
"org.opencontainers.image.version": "{{.Version}}"
|
||||||
|
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
||||||
|
"org.opencontainers.image.source": "{{.GitURL}}"
|
||||||
|
annotations:
|
||||||
|
"org.opencontainers.image.created": "{{.Date}}"
|
||||||
|
"org.opencontainers.image.title": "{{.ProjectName}}"
|
||||||
|
"org.opencontainers.image.version": "{{.Version}}"
|
||||||
|
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
||||||
|
"org.opencontainers.image.source": "{{.GitURL}}"
|
||||||
|
"maintainer": "dev@netbird.io"
|
||||||
- id: relay
|
- id: relay
|
||||||
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
|
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
|
||||||
ids:
|
ids:
|
||||||
@@ -365,7 +477,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
|
||||||
@@ -421,6 +533,41 @@ dockers_v2:
|
|||||||
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
||||||
"org.opencontainers.image.source": "{{.GitURL}}"
|
"org.opencontainers.image.source": "{{.GitURL}}"
|
||||||
"maintainer": "dev@netbird.io"
|
"maintainer": "dev@netbird.io"
|
||||||
|
- id: proxy-ubi
|
||||||
|
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
|
||||||
|
ids:
|
||||||
|
- netbird-proxy
|
||||||
|
images:
|
||||||
|
- netbirdio/reverse-proxy
|
||||||
|
- ghcr.io/netbirdio/reverse-proxy
|
||||||
|
tags:
|
||||||
|
- "{{ .Version }}-ubi"
|
||||||
|
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}ubi-latest{{ end }}"
|
||||||
|
dockerfile: proxy/Dockerfile.ubi
|
||||||
|
platforms:
|
||||||
|
- linux/amd64
|
||||||
|
- linux/arm64
|
||||||
|
build_args:
|
||||||
|
VERSION: "{{ .Version }}"
|
||||||
|
RELEASE: "{{ .Timestamp }}"
|
||||||
|
hooks:
|
||||||
|
pre:
|
||||||
|
- cmd: 'sh proxy/collect-licenses.sh "{{ .ContextDir }}/licenses" amd64 arm64'
|
||||||
|
env:
|
||||||
|
- GOOS=linux
|
||||||
|
- CGO_ENABLED=0
|
||||||
|
labels:
|
||||||
|
"org.opencontainers.image.created": "{{.Date}}"
|
||||||
|
"org.opencontainers.image.version": "{{.Version}}"
|
||||||
|
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
||||||
|
"org.opencontainers.image.source": "{{.GitURL}}"
|
||||||
|
annotations:
|
||||||
|
"org.opencontainers.image.created": "{{.Date}}"
|
||||||
|
"org.opencontainers.image.title": "{{.ProjectName}}"
|
||||||
|
"org.opencontainers.image.version": "{{.Version}}"
|
||||||
|
"org.opencontainers.image.revision": "{{.FullCommit}}"
|
||||||
|
"org.opencontainers.image.source": "{{.GitURL}}"
|
||||||
|
"maintainer": "dev@netbird.io"
|
||||||
|
|
||||||
brews:
|
brews:
|
||||||
- ids:
|
- ids:
|
||||||
@@ -453,7 +600,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
|
||||||
|
|||||||
+2
-2
@@ -192,7 +192,7 @@ dependencies are installed. Here is a short guide on how that can be done.
|
|||||||
|
|
||||||
### Requirements
|
### Requirements
|
||||||
|
|
||||||
#### Go 1.25
|
#### Go 1.26
|
||||||
|
|
||||||
Follow the installation guide from https://go.dev/
|
Follow the installation guide from https://go.dev/
|
||||||
|
|
||||||
@@ -200,7 +200,7 @@ Follow the installation guide from https://go.dev/
|
|||||||
|
|
||||||
The desktop UI client (`client/ui`) is built with [Wails v3](https://v3.wails.io/) and a React frontend rendered in a WebView. To build it you need:
|
The desktop UI client (`client/ui`) is built with [Wails v3](https://v3.wails.io/) and a React frontend rendered in a WebView. To build it you need:
|
||||||
|
|
||||||
- Go ≥ 1.25
|
- Go ≥ 1.26
|
||||||
- Node ≥ 20 and **pnpm** (`corepack enable && corepack prepare pnpm@latest --activate`)
|
- Node ≥ 20 and **pnpm** (`corepack enable && corepack prepare pnpm@latest --activate`)
|
||||||
- The `wails3` CLI: `go install github.com/wailsapp/wails/v3/cmd/wails3@latest`
|
- The `wails3` CLI: `go install github.com/wailsapp/wails/v3/cmd/wails3@latest`
|
||||||
- The `task` runner: `go install github.com/go-task/task/v3/cmd/task@latest`
|
- The `task` runner: `go install github.com/go-task/task/v3/cmd/task@latest`
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -26,7 +26,7 @@
|
|||||||
<strong>
|
<strong>
|
||||||
Start using NetBird at <a href="https://netbird.io/pricing">netbird.io</a>
|
Start using NetBird at <a href="https://netbird.io/pricing">netbird.io</a>
|
||||||
<br/>
|
<br/>
|
||||||
See <a href="https://netbird.io/docs/">Documentation</a>
|
See <a href="https://docs.netbird.io/">Documentation</a>
|
||||||
<br/>
|
<br/>
|
||||||
Join our <a href="https://docs.netbird.io/slack-url">Slack channel</a> or our <a href="https://forum.netbird.io">Community forum</a>
|
Join our <a href="https://docs.netbird.io/slack-url">Slack channel</a> or our <a href="https://forum.netbird.io">Community forum</a>
|
||||||
</strong>
|
</strong>
|
||||||
@@ -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)
|
||||||
@@ -96,6 +96,42 @@ components:
|
|||||||
— the management-side control plane: providers, policies, guardrails, limits, routing,
|
— the management-side control plane: providers, policies, guardrails, limits, routing,
|
||||||
and usage/access logs.
|
and usage/access logs.
|
||||||
|
|
||||||
|
## Access roles
|
||||||
|
|
||||||
|
Agent Network permissions build on the account permission matrix
|
||||||
|
([`management/server/permissions/`](../management/server/permissions)). The
|
||||||
|
`agent_network` area is split into dotted submodules (`agent_network.providers`,
|
||||||
|
`.policies`, `.guardrails`, `.budgets`, `.usage`, `.logs`, `.settings`); a role may
|
||||||
|
grant a single submodule or the parent, which cascades to all of them.
|
||||||
|
|
||||||
|
Two roles delegate Agent Network access without account-admin rights:
|
||||||
|
|
||||||
|
- **`agent_network_admin`** — full control over the whole `agent_network` area plus
|
||||||
|
read-only users, groups, peers, and account info (needed to build policies).
|
||||||
|
Nothing else in the account.
|
||||||
|
- **`usage_viewer`** — the regular User baseline plus read on
|
||||||
|
`agent_network.usage` (the aggregated usage and cost overview) and
|
||||||
|
`agent_network.logs` (the account-wide request-level access logs, which can
|
||||||
|
contain captured prompts), and read-only access to the resources those
|
||||||
|
filters resolve against: users, groups, peers, and the provider list
|
||||||
|
(connection config redacted — no upstream URLs or operator-supplied header
|
||||||
|
values). No policies, guardrails, budgets, or settings.
|
||||||
|
|
||||||
|
Every authenticated user, regardless of role, can read the caller-scoped
|
||||||
|
self-service endpoint `GET /api/agent-network/agent-config` (the endpoint, providers,
|
||||||
|
and models the caller's own policies allow — what a local AI tool needs and nothing
|
||||||
|
more). The regular usage and access-log endpoints self-scope instead of denying:
|
||||||
|
a caller without the account-wide grant gets their own rows back, so "my usage"
|
||||||
|
and "my requests" are the same endpoints the admin dashboard uses. The provider
|
||||||
|
list self-scopes the same way — a caller without the providers grant gets the
|
||||||
|
providers their own policies authorize, reduced to the display surface, with
|
||||||
|
each provider's model list cut to what the caller's policy guardrails and the
|
||||||
|
provider's declared models effectively permit (the same computation the setup
|
||||||
|
answer and the proxy use). This feeds the dashboard's provider and model
|
||||||
|
filters. Role
|
||||||
|
definitions live in
|
||||||
|
[`management/server/permissions/roles/`](../management/server/permissions/roles).
|
||||||
|
|
||||||
## Documentation
|
## Documentation
|
||||||
|
|
||||||
Full documentation, architecture, and quickstart:
|
Full documentation, architecture, and quickstart:
|
||||||
|
|||||||
+52
-33
@@ -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"]
|
||||||
@@ -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,6 +91,14 @@ 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
|
||||||
|
|
||||||
@@ -178,6 +187,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 +213,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 +240,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 +257,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 +329,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 +353,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 +399,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 +498,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 +512,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 +521,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 +552,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
|
||||||
|
}
|
||||||
+17
-16
@@ -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)
|
|
||||||
return true, err
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// LoginWithSetupKeyAndSaveConfig test the connectivity with the management server with the setup key.
|
// LoginWithSetupKeyAndSaveConfig registers the peer with the setup key; the config is already persisted by NewAuth.
|
||||||
func (a *Auth) LoginWithSetupKeyAndSaveConfig(resultListener ErrListener, setupKey string, deviceName string) {
|
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
|
||||||
|
|||||||
@@ -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,11 +18,30 @@ 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
|
||||||
}
|
}
|
||||||
@@ -27,7 +50,7 @@ func (p *Preferences) GetManagementURL() (string, error) {
|
|||||||
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
|
||||||
@@ -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.ReadConfig(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,6 +105,9 @@ 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
|
||||||
}
|
}
|
||||||
@@ -96,6 +126,9 @@ 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
|
||||||
}
|
}
|
||||||
@@ -109,6 +142,9 @@ 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
|
||||||
}
|
}
|
||||||
@@ -127,6 +163,9 @@ 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
|
||||||
}
|
}
|
||||||
@@ -181,6 +220,9 @@ 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
|
||||||
}
|
}
|
||||||
@@ -291,6 +333,9 @@ 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
|
||||||
}
|
}
|
||||||
@@ -325,8 +370,34 @@ func (p *Preferences) SetDisableIPv6(disable bool) {
|
|||||||
p.configInput.DisableIPv6 = &disable
|
p.configInput.DisableIPv6 = &disable
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetRemoteJobsAllowed reads the remote jobs opt-in from config file
|
||||||
|
func (p *Preferences) GetRemoteJobsAllowed() (bool, error) {
|
||||||
|
policy := p.policy()
|
||||||
|
if !policy.HasKey(mdm.KeyRemoteJobsAllowed) && p.configInput.RemoteJobsAllowed != nil {
|
||||||
|
return *p.configInput.RemoteJobsAllowed, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
cfg.ApplyMDMPolicy(policy)
|
||||||
|
if cfg.RemoteJobsAllowed == nil {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
return *cfg.RemoteJobsAllowed, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetRemoteJobsAllowed stores the given value and waits for commit
|
||||||
|
func (p *Preferences) SetRemoteJobsAllowed(allowed bool) {
|
||||||
|
p.configInput.RemoteJobsAllowed = &allowed
|
||||||
|
}
|
||||||
|
|
||||||
// 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) {
|
||||||
|
|||||||
@@ -0,0 +1,125 @@
|
|||||||
|
package android
|
||||||
|
|
||||||
|
// SplitTunnelMode is which of the two selections, if either, the tunnel applies.
|
||||||
|
// Its values land in the profile's stored preferences, so the constants below
|
||||||
|
// are append-only and must never be reordered.
|
||||||
|
type SplitTunnelMode int
|
||||||
|
|
||||||
|
const (
|
||||||
|
modeOff SplitTunnelMode = iota
|
||||||
|
modeExclude
|
||||||
|
modeInclude
|
||||||
|
)
|
||||||
|
|
||||||
|
// The same modes as basic ints. gomobile drops a constant whose type is not a
|
||||||
|
// basic one, so these are what reaches the generated Java bindings, and they
|
||||||
|
// keep the Android side tied to the values above instead of repeating 0, 1, 2.
|
||||||
|
const (
|
||||||
|
SplitTunnelModeOff = int(modeOff)
|
||||||
|
SplitTunnelModeExclude = int(modeExclude)
|
||||||
|
SplitTunnelModeInclude = int(modeInclude)
|
||||||
|
)
|
||||||
|
|
||||||
|
type splitTunnelSection struct {
|
||||||
|
Mode SplitTunnelMode `json:"mode"`
|
||||||
|
Excluded []string `json:"excluded"`
|
||||||
|
Included []string `json:"included"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PackageList wraps []string for gomobile compatibility.
|
||||||
|
type PackageList struct {
|
||||||
|
items []string
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewPackageList creates an empty list to fill via Add.
|
||||||
|
func NewPackageList() *PackageList {
|
||||||
|
return &PackageList{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add appends a package name, ignoring empty ones.
|
||||||
|
func (l *PackageList) Add(s string) {
|
||||||
|
if s == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
l.items = append(l.items, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Size returns the number of entries.
|
||||||
|
func (l *PackageList) Size() int {
|
||||||
|
return len(l.items)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get returns the entry at index i, or an empty string when out of range.
|
||||||
|
func (l *PackageList) Get(i int) string {
|
||||||
|
if i < 0 || i >= len(l.items) {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return l.items[i]
|
||||||
|
}
|
||||||
|
|
||||||
|
// SplitTunnelSettings is one profile's choice of which applications the tunnel
|
||||||
|
// carries. The two selections are kept apart because the platform applies one
|
||||||
|
// or the other and never both, and so that switching mode does not throw away
|
||||||
|
// the picks made in the other one.
|
||||||
|
//
|
||||||
|
// Mode is an int rather than a SplitTunnelMode because gomobile carries only
|
||||||
|
// basic types across the binding. It holds one of the SplitTunnelMode*
|
||||||
|
// constants.
|
||||||
|
type SplitTunnelSettings struct {
|
||||||
|
Mode int
|
||||||
|
Excluded *PackageList
|
||||||
|
Included *PackageList
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewSplitTunnelSettings creates settings that carry every application.
|
||||||
|
func NewSplitTunnelSettings() *SplitTunnelSettings {
|
||||||
|
return &SplitTunnelSettings{
|
||||||
|
Mode: SplitTunnelModeOff,
|
||||||
|
Excluded: NewPackageList(),
|
||||||
|
Included: NewPackageList(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func packagesOf(list *PackageList) []string {
|
||||||
|
if list == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
out := make([]string, 0, len(list.items))
|
||||||
|
out = append(out, list.items...)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// normalizeSplitTunnelMode maps anything outside the known set to off, so a mode
|
||||||
|
// written by a newer build degrades to carrying every application rather than to
|
||||||
|
// some other mode's behaviour.
|
||||||
|
func normalizeSplitTunnelMode(mode SplitTunnelMode) SplitTunnelMode {
|
||||||
|
switch mode {
|
||||||
|
case modeExclude, modeInclude:
|
||||||
|
return mode
|
||||||
|
default:
|
||||||
|
return modeOff
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func settingsFromSection(section splitTunnelSection) *SplitTunnelSettings {
|
||||||
|
out := NewSplitTunnelSettings()
|
||||||
|
out.Mode = int(normalizeSplitTunnelMode(section.Mode))
|
||||||
|
for _, pkg := range section.Excluded {
|
||||||
|
out.Excluded.Add(pkg)
|
||||||
|
}
|
||||||
|
for _, pkg := range section.Included {
|
||||||
|
out.Included.Add(pkg)
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func sectionFromSettings(settings *SplitTunnelSettings) splitTunnelSection {
|
||||||
|
if settings == nil {
|
||||||
|
settings = NewSplitTunnelSettings()
|
||||||
|
}
|
||||||
|
return splitTunnelSection{
|
||||||
|
Mode: normalizeSplitTunnelMode(SplitTunnelMode(settings.Mode)),
|
||||||
|
Excluded: packagesOf(settings.Excluded),
|
||||||
|
Included: packagesOf(settings.Included),
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,34 @@
|
|||||||
|
//go:build android
|
||||||
|
|
||||||
|
package android
|
||||||
|
|
||||||
|
const splitTunnelNamespace = "split-tunnel"
|
||||||
|
|
||||||
|
// SplitTunnelStore reads and writes a profile's split tunnelling settings.
|
||||||
|
type SplitTunnelStore struct {
|
||||||
|
prefs prefsStore
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewSplitTunnelStore opens the split tunnelling store of the given profile.
|
||||||
|
func NewSplitTunnelStore(configDir, profileID string) (*SplitTunnelStore, error) {
|
||||||
|
prefs, err := newProfilePrefs(configDir, profileID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &SplitTunnelStore{prefs: prefs}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load returns the stored settings, or settings that carry every application
|
||||||
|
// when the profile has none saved.
|
||||||
|
func (s *SplitTunnelStore) Load() (*SplitTunnelSettings, error) {
|
||||||
|
var section splitTunnelSection
|
||||||
|
if _, err := s.prefs.Get(splitTunnelNamespace, §ion); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return settingsFromSection(section), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Save replaces the stored settings.
|
||||||
|
func (s *SplitTunnelStore) Save(settings *SplitTunnelSettings) error {
|
||||||
|
return s.prefs.Put(splitTunnelNamespace, sectionFromSettings(settings))
|
||||||
|
}
|
||||||
@@ -0,0 +1,151 @@
|
|||||||
|
package android
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNormalizeSplitTunnelMode(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
mode SplitTunnelMode
|
||||||
|
want SplitTunnelMode
|
||||||
|
}{
|
||||||
|
{name: "exclude is kept", mode: modeExclude, want: modeExclude},
|
||||||
|
{name: "include is kept", mode: modeInclude, want: modeInclude},
|
||||||
|
{name: "off is kept", mode: modeOff, want: modeOff},
|
||||||
|
{name: "a mode from a newer build falls back to off", mode: SplitTunnelMode(7), want: modeOff},
|
||||||
|
{name: "a negative mode falls back to off", mode: SplitTunnelMode(-1), want: modeOff},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := normalizeSplitTunnelMode(tt.mode); got != tt.want {
|
||||||
|
t.Errorf("normalizeSplitTunnelMode(%d) = %d, want %d", tt.mode, got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The constants the Android side reads must stay the values the store writes:
|
||||||
|
// gomobile carries the ints below, not the typed constants they mirror.
|
||||||
|
func TestSplitTunnelModeConstantsMirrorTheTypedOnes(t *testing.T) {
|
||||||
|
if SplitTunnelModeOff != int(modeOff) {
|
||||||
|
t.Errorf("off = %d, want %d", SplitTunnelModeOff, modeOff)
|
||||||
|
}
|
||||||
|
if SplitTunnelModeExclude != int(modeExclude) {
|
||||||
|
t.Errorf("exclude = %d, want %d", SplitTunnelModeExclude, modeExclude)
|
||||||
|
}
|
||||||
|
if SplitTunnelModeInclude != int(modeInclude) {
|
||||||
|
t.Errorf("include = %d, want %d", SplitTunnelModeInclude, modeInclude)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSettingsFromSection(t *testing.T) {
|
||||||
|
got := settingsFromSection(splitTunnelSection{
|
||||||
|
Mode: modeExclude,
|
||||||
|
Excluded: []string{"com.example.a", "com.example.b"},
|
||||||
|
Included: []string{"com.example.c"},
|
||||||
|
})
|
||||||
|
|
||||||
|
if got.Mode != SplitTunnelModeExclude {
|
||||||
|
t.Errorf("mode = %d, want %d", got.Mode, SplitTunnelModeExclude)
|
||||||
|
}
|
||||||
|
if got.Excluded.Size() != 2 || got.Excluded.Get(0) != "com.example.a" {
|
||||||
|
t.Errorf("excluded = %v, want the two stored packages", packagesOf(got.Excluded))
|
||||||
|
}
|
||||||
|
if got.Included.Size() != 1 || got.Included.Get(0) != "com.example.c" {
|
||||||
|
t.Errorf("included = %v, want the stored package", packagesOf(got.Included))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A profile that has never stored anything decodes into an empty section, and
|
||||||
|
// must come back as settings that carry every application rather than as nil
|
||||||
|
// lists the caller would have to guard against.
|
||||||
|
func TestSettingsFromEmptySectionCarriesEverything(t *testing.T) {
|
||||||
|
got := settingsFromSection(splitTunnelSection{})
|
||||||
|
|
||||||
|
if got.Mode != SplitTunnelModeOff {
|
||||||
|
t.Errorf("mode = %d, want %d", got.Mode, SplitTunnelModeOff)
|
||||||
|
}
|
||||||
|
if got.Excluded == nil || got.Included == nil {
|
||||||
|
t.Fatal("both selections must be usable lists, not nil")
|
||||||
|
}
|
||||||
|
if got.Excluded.Size() != 0 || got.Included.Size() != 0 {
|
||||||
|
t.Errorf("selections = %v/%v, want both empty", packagesOf(got.Excluded), packagesOf(got.Included))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The section is what the profile's preference file holds, so the mode has to
|
||||||
|
// survive a JSON round trip as the number the constants name.
|
||||||
|
func TestSectionEncodesTheModeAsItsNumber(t *testing.T) {
|
||||||
|
raw, err := json.Marshal(sectionFromSettings(&SplitTunnelSettings{Mode: SplitTunnelModeInclude}))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marshal section: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var back splitTunnelSection
|
||||||
|
if err := json.Unmarshal(raw, &back); err != nil {
|
||||||
|
t.Fatalf("unmarshal section: %v", err)
|
||||||
|
}
|
||||||
|
if back.Mode != modeInclude {
|
||||||
|
t.Errorf("mode = %d, want %d, from %s", back.Mode, modeInclude, raw)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSectionFromSettingsRoundTrip(t *testing.T) {
|
||||||
|
settings := NewSplitTunnelSettings()
|
||||||
|
settings.Mode = SplitTunnelModeInclude
|
||||||
|
settings.Included.Add("com.example.a")
|
||||||
|
settings.Excluded.Add("com.example.b")
|
||||||
|
|
||||||
|
section := sectionFromSettings(settings)
|
||||||
|
back := settingsFromSection(section)
|
||||||
|
|
||||||
|
if back.Mode != SplitTunnelModeInclude {
|
||||||
|
t.Errorf("mode = %d, want %d", back.Mode, SplitTunnelModeInclude)
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(packagesOf(back.Included), []string{"com.example.a"}) {
|
||||||
|
t.Errorf("included = %v, want [com.example.a]", packagesOf(back.Included))
|
||||||
|
}
|
||||||
|
// The inactive selection survives, so switching mode back does not make the
|
||||||
|
// user pick their applications again.
|
||||||
|
if !reflect.DeepEqual(packagesOf(back.Excluded), []string{"com.example.b"}) {
|
||||||
|
t.Errorf("excluded = %v, want [com.example.b]", packagesOf(back.Excluded))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A mode the Java side never sets, such as one left by a newer build, must not
|
||||||
|
// reach the stored section either.
|
||||||
|
func TestSectionFromSettingsNormalizesAnUnknownMode(t *testing.T) {
|
||||||
|
section := sectionFromSettings(&SplitTunnelSettings{Mode: 7})
|
||||||
|
|
||||||
|
if section.Mode != modeOff {
|
||||||
|
t.Errorf("mode = %d, want %d", section.Mode, modeOff)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSectionFromNilSettings(t *testing.T) {
|
||||||
|
section := sectionFromSettings(nil)
|
||||||
|
|
||||||
|
if section.Mode != modeOff {
|
||||||
|
t.Errorf("mode = %d, want %d", section.Mode, modeOff)
|
||||||
|
}
|
||||||
|
if len(section.Excluded) != 0 || len(section.Included) != 0 {
|
||||||
|
t.Errorf("selections = %v/%v, want both empty", section.Excluded, section.Included)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPackageListIgnoresEmptyAndBounds(t *testing.T) {
|
||||||
|
list := NewPackageList()
|
||||||
|
list.Add("com.example.a")
|
||||||
|
list.Add("")
|
||||||
|
|
||||||
|
if list.Size() != 1 {
|
||||||
|
t.Errorf("size = %d, want 1", list.Size())
|
||||||
|
}
|
||||||
|
if list.Get(-1) != "" || list.Get(5) != "" {
|
||||||
|
t.Error("out of range access must return an empty string")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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,
|
||||||
|
|||||||
+59
-28
@@ -3,7 +3,6 @@ package cmd
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os/user"
|
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -24,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
|
||||||
@@ -114,7 +116,7 @@ func debugConfigDump(cmd *cobra.Command, _ []string) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get active profile: %v", err)
|
return fmt.Errorf("get active profile: %v", err)
|
||||||
}
|
}
|
||||||
currUser, err := user.Current()
|
currUser, err := profilemanager.InvokingUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get current user: %v", err)
|
return fmt.Errorf("get current user: %v", err)
|
||||||
}
|
}
|
||||||
@@ -258,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 {
|
||||||
@@ -285,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() {
|
||||||
@@ -402,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 {
|
||||||
@@ -459,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()
|
||||||
@@ -547,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")
|
||||||
|
}
|
||||||
@@ -0,0 +1,13 @@
|
|||||||
|
package cmd
|
||||||
|
|
||||||
|
// remoteJobsAllowedFlag opts this peer into running remote jobs (e.g. debug
|
||||||
|
// bundles) requested by the management server. It defaults to false: remote
|
||||||
|
// jobs are an explicit opt-in, and enabling it is a privileged change (see the
|
||||||
|
// daemon gate in client/server), mirroring the SSH server opt-in.
|
||||||
|
const remoteJobsAllowedFlag = "allow-remote-jobs"
|
||||||
|
|
||||||
|
var remoteJobsAllowed bool
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
upCmd.PersistentFlags().BoolVar(&remoteJobsAllowed, remoteJobsAllowedFlag, false, "Allow the management server to run remote jobs (e.g. debug bundles) on this peer")
|
||||||
|
}
|
||||||
+24
-20
@@ -4,7 +4,6 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"os/user"
|
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
@@ -16,6 +15,7 @@ import (
|
|||||||
"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"
|
||||||
@@ -53,7 +53,7 @@ var loginCmd = &cobra.Command{
|
|||||||
// nolint
|
// nolint
|
||||||
ctx = context.WithValue(ctx, system.DeviceNameCtxKey, hostName)
|
ctx = context.WithValue(ctx, system.DeviceNameCtxKey, hostName)
|
||||||
}
|
}
|
||||||
username, err := user.Current()
|
username, err := profilemanager.InvokingUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get current user: %v", err)
|
return fmt.Errorf("get current user: %v", err)
|
||||||
}
|
}
|
||||||
@@ -74,7 +74,7 @@ var loginCmd = &cobra.Command{
|
|||||||
if providedSetupKey != "" {
|
if providedSetupKey != "" {
|
||||||
return fmt.Errorf("--extend cannot be combined with a setup key; setup keys can only enrol new peers")
|
return fmt.Errorf("--extend cannot be combined with a setup key; setup keys can only enrol new peers")
|
||||||
}
|
}
|
||||||
if err := doExtendSession(ctx, cmd); err != nil {
|
if err := doExtendSession(ctx, cmd, activeProf); err != nil {
|
||||||
return fmt.Errorf("extend session failed: %v", err)
|
return fmt.Errorf("extend session failed: %v", err)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
@@ -92,7 +92,7 @@ var loginCmd = &cobra.Command{
|
|||||||
return fmt.Errorf("daemon login failed: %v", err)
|
return fmt.Errorf("daemon login failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
cmd.Println("Logging successfully")
|
cmd.Println("Login successful")
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
},
|
},
|
||||||
@@ -176,7 +176,7 @@ func doDaemonLogin(ctx context.Context, cmd *cobra.Command, providedSetupKey str
|
|||||||
// (browser + verification URL) and the resulting JWT is forwarded to the
|
// (browser + verification URL) and the resulting JWT is forwarded to the
|
||||||
// management server's ExtendAuthSession RPC. The tunnel stays up
|
// management server's ExtendAuthSession RPC. The tunnel stays up
|
||||||
// throughout — no Down/Up, no network-map resync.
|
// throughout — no Down/Up, no network-map resync.
|
||||||
func doExtendSession(ctx context.Context, cmd *cobra.Command) error {
|
func doExtendSession(ctx context.Context, cmd *cobra.Command, activeProf *profilemanager.Profile) error {
|
||||||
conn, err := DialClientGRPCServer(ctx, daemonAddr)
|
conn, err := DialClientGRPCServer(ctx, daemonAddr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
//nolint
|
//nolint
|
||||||
@@ -190,14 +190,12 @@ func doExtendSession(ctx context.Context, cmd *cobra.Command) error {
|
|||||||
|
|
||||||
// the CLI runs in the user's session, the daemon does not: tell it what we can see
|
// the CLI runs in the user's session, the daemon does not: tell it what we can see
|
||||||
req := &proto.RequestExtendAuthSessionRequest{HasGraphicalSession: util.HasGraphicalSession()}
|
req := &proto.RequestExtendAuthSessionRequest{HasGraphicalSession: util.HasGraphicalSession()}
|
||||||
// Pre-fill the IdP login hint from the active profile so the user
|
// Pre-fill the IdP login hint from the resolved profile so the user
|
||||||
// doesn't have to retype their email. Best-effort: we still proceed
|
// doesn't have to retype their email. Best-effort: we still proceed
|
||||||
// without a hint if the lookup fails.
|
// without a hint if the lookup fails.
|
||||||
pm := profilemanager.NewProfileManager()
|
pm := profilemanager.NewProfileManager()
|
||||||
if active, perr := pm.GetActiveProfile(); perr == nil {
|
if profState, perr := pm.GetProfileState(activeProf.ID); perr == nil && profState.Email != "" {
|
||||||
if profState, sperr := pm.GetProfileState(active.ID); sperr == nil && profState.Email != "" {
|
req.Hint = &profState.Email
|
||||||
req.Hint = &profState.Email
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
startResp, err := client.RequestExtendAuthSession(ctx, req)
|
startResp, err := client.RequestExtendAuthSession(ctx, req)
|
||||||
@@ -235,9 +233,11 @@ func getActiveProfile(ctx context.Context, pm *profilemanager.ProfileManager, pr
|
|||||||
// switch profile if provided
|
// switch profile if provided
|
||||||
|
|
||||||
if profileName != "" {
|
if profileName != "" {
|
||||||
if err := switchProfileOnDaemon(ctx, pm, profileName, username); err != nil {
|
prof, err := switchProfileOnDaemon(ctx, pm, profileName, username)
|
||||||
|
if err != nil {
|
||||||
return nil, fmt.Errorf("switch profile: %v", err)
|
return nil, fmt.Errorf("switch profile: %v", err)
|
||||||
}
|
}
|
||||||
|
return prof, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
activeProf, err := pm.GetActiveProfile()
|
activeProf, err := pm.GetActiveProfile()
|
||||||
@@ -251,20 +251,19 @@ func getActiveProfile(ctx context.Context, pm *profilemanager.ProfileManager, pr
|
|||||||
return activeProf, nil
|
return activeProf, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func switchProfileOnDaemon(ctx context.Context, pm *profilemanager.ProfileManager, handle string, username string) error {
|
func switchProfileOnDaemon(ctx context.Context, pm *profilemanager.ProfileManager, handle string, username string) (*profilemanager.Profile, error) {
|
||||||
resolvedID, err := switchProfile(ctx, handle, username)
|
resolvedID, err := switchProfile(ctx, handle, username)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("switch profile on daemon: %v", err)
|
return nil, fmt.Errorf("switch profile on daemon: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := pm.SwitchProfile(resolvedID); err != nil {
|
if err := pm.SwitchProfile(resolvedID); err != nil {
|
||||||
return fmt.Errorf("switch profile: %v", err)
|
return nil, fmt.Errorf("switch profile: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
conn, err := DialClientGRPCServer(ctx, daemonAddr)
|
conn, err := DialClientGRPCServer(ctx, daemonAddr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Errorf("failed to connect to service CLI interface %v", err)
|
return nil, fmt.Errorf("connect to service CLI interface: %w", err)
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
|
|
||||||
@@ -272,17 +271,17 @@ func switchProfileOnDaemon(ctx context.Context, pm *profilemanager.ProfileManage
|
|||||||
|
|
||||||
status, err := client.Status(ctx, &proto.StatusRequest{})
|
status, err := client.Status(ctx, &proto.StatusRequest{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("unable to get daemon status: %v", err)
|
return nil, fmt.Errorf("unable to get daemon status: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if status.Status == string(internal.StatusConnected) {
|
if status.Status == string(internal.StatusConnected) {
|
||||||
if _, err := client.Down(ctx, &proto.DownRequest{}); err != nil {
|
if _, err := client.Down(ctx, &proto.DownRequest{}); err != nil {
|
||||||
log.Errorf("call service down method: %v", err)
|
log.Errorf("call service down method: %v", err)
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return &profilemanager.Profile{ID: resolvedID}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// switchProfile asks the daemon to switch to the profile identified by
|
// switchProfile asks the daemon to switch to the profile identified by
|
||||||
@@ -332,6 +331,11 @@ func doForegroundLogin(ctx context.Context, cmd *cobra.Command, setupKey string,
|
|||||||
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)
|
||||||
}
|
}
|
||||||
|
// 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
|
||||||
@@ -345,7 +349,7 @@ func doForegroundLogin(ctx context.Context, cmd *cobra.Command, setupKey string,
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("foreground login failed: %v", err)
|
return fmt.Errorf("foreground login failed: %v", err)
|
||||||
}
|
}
|
||||||
cmd.Println("Logging successfully")
|
cmd.Println("Login successful")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -3,11 +3,11 @@ package cmd
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os/user"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
"github.com/netbirdio/netbird/client/proto"
|
"github.com/netbirdio/netbird/client/proto"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -37,7 +37,7 @@ var logoutCmd = &cobra.Command{
|
|||||||
if profileName != "" {
|
if profileName != "" {
|
||||||
req.ProfileName = &profileName
|
req.ProfileName = &profileName
|
||||||
|
|
||||||
currUser, err := user.Current()
|
currUser, err := profilemanager.InvokingUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get current user: %v", err)
|
return fmt.Errorf("get current user: %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os/user"
|
|
||||||
"strings"
|
"strings"
|
||||||
"text/tabwriter"
|
"text/tabwriter"
|
||||||
"time"
|
"time"
|
||||||
@@ -97,7 +96,7 @@ func listProfilesFunc(cmd *cobra.Command, _ []string) error {
|
|||||||
}
|
}
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
|
|
||||||
currUser, err := user.Current()
|
currUser, err := profilemanager.InvokingUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get current user: %w", err)
|
return fmt.Errorf("get current user: %w", err)
|
||||||
}
|
}
|
||||||
@@ -138,7 +137,7 @@ func addProfileFunc(cmd *cobra.Command, args []string) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
currUser, err := user.Current()
|
currUser, err := profilemanager.InvokingUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get current user: %w", err)
|
return fmt.Errorf("get current user: %w", err)
|
||||||
}
|
}
|
||||||
@@ -179,7 +178,7 @@ func renameProfileFunc(cmd *cobra.Command, args []string) error {
|
|||||||
}
|
}
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
|
|
||||||
currUser, err := user.Current()
|
currUser, err := profilemanager.InvokingUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get current user: %w", err)
|
return fmt.Errorf("get current user: %w", err)
|
||||||
}
|
}
|
||||||
@@ -233,7 +232,7 @@ func removeProfileFunc(cmd *cobra.Command, args []string) error {
|
|||||||
}
|
}
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
|
|
||||||
currUser, err := user.Current()
|
currUser, err := profilemanager.InvokingUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get current user: %w", err)
|
return fmt.Errorf("get current user: %w", err)
|
||||||
}
|
}
|
||||||
@@ -261,7 +260,7 @@ func selectProfileFunc(cmd *cobra.Command, args []string) error {
|
|||||||
profileManager := profilemanager.NewProfileManager()
|
profileManager := profilemanager.NewProfileManager()
|
||||||
handle := args[0]
|
handle := args[0]
|
||||||
|
|
||||||
currUser, err := user.Current()
|
currUser, err := profilemanager.InvokingUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get current user: %w", err)
|
return fmt.Errorf("get current user: %w", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ import (
|
|||||||
|
|
||||||
"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"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/localmetrics"
|
||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -31,6 +32,8 @@ const (
|
|||||||
dnsResolverAddress = "dns-resolver-address"
|
dnsResolverAddress = "dns-resolver-address"
|
||||||
enableRosenpassFlag = "enable-rosenpass"
|
enableRosenpassFlag = "enable-rosenpass"
|
||||||
rosenpassPermissiveFlag = "rosenpass-permissive"
|
rosenpassPermissiveFlag = "rosenpass-permissive"
|
||||||
|
enableLocalMetricsFlag = "enable-local-metrics"
|
||||||
|
localMetricsAddressFlag = "local-metrics-address"
|
||||||
preSharedKeyFlag = "preshared-key"
|
preSharedKeyFlag = "preshared-key"
|
||||||
interfaceNameFlag = "interface-name"
|
interfaceNameFlag = "interface-name"
|
||||||
wireguardPortFlag = "wireguard-port"
|
wireguardPortFlag = "wireguard-port"
|
||||||
@@ -80,6 +83,8 @@ var (
|
|||||||
updateSettingsDisabled bool
|
updateSettingsDisabled bool
|
||||||
captureEnabled bool
|
captureEnabled bool
|
||||||
networksDisabled bool
|
networksDisabled bool
|
||||||
|
localMetricsEnabled bool
|
||||||
|
localMetricsAddr string
|
||||||
|
|
||||||
rootCmd = &cobra.Command{
|
rootCmd = &cobra.Command{
|
||||||
Use: "netbird",
|
Use: "netbird",
|
||||||
@@ -215,6 +220,8 @@ func init() {
|
|||||||
upCmd.PersistentFlags().BoolVar(&rosenpassEnabled, enableRosenpassFlag, false, "[Experimental] Enable Rosenpass feature. If enabled, the connection will be post-quantum secured via Rosenpass.")
|
upCmd.PersistentFlags().BoolVar(&rosenpassEnabled, enableRosenpassFlag, false, "[Experimental] Enable Rosenpass feature. If enabled, the connection will be post-quantum secured via Rosenpass.")
|
||||||
upCmd.PersistentFlags().BoolVar(&rosenpassPermissive, rosenpassPermissiveFlag, false, "[Experimental] Enable Rosenpass in permissive mode to allow this peer to accept WireGuard connections without requiring Rosenpass functionality from peers that do not have Rosenpass enabled.")
|
upCmd.PersistentFlags().BoolVar(&rosenpassPermissive, rosenpassPermissiveFlag, false, "[Experimental] Enable Rosenpass in permissive mode to allow this peer to accept WireGuard connections without requiring Rosenpass functionality from peers that do not have Rosenpass enabled.")
|
||||||
upCmd.PersistentFlags().BoolVar(&autoConnectDisabled, disableAutoConnectFlag, false, "Disables auto-connect feature. If enabled, then the client won't connect automatically when the service starts.")
|
upCmd.PersistentFlags().BoolVar(&autoConnectDisabled, disableAutoConnectFlag, false, "Disables auto-connect feature. If enabled, then the client won't connect automatically when the service starts.")
|
||||||
|
upCmd.PersistentFlags().BoolVar(&localMetricsEnabled, enableLocalMetricsFlag, false, "Enables a local Prometheus /metrics endpoint exposing connection state (peers, latency, P2P vs relay).")
|
||||||
|
upCmd.PersistentFlags().StringVar(&localMetricsAddr, localMetricsAddressFlag, localmetrics.DefaultListenAddress, "Listen address of the local Prometheus /metrics endpoint.")
|
||||||
upCmd.PersistentFlags().BoolVar(&lazyConnEnabled, enableLazyConnectionFlag, false, "Deprecated: no longer used. Lazy connections are controlled by the server and the NB_LAZY_CONN environment variable.")
|
upCmd.PersistentFlags().BoolVar(&lazyConnEnabled, enableLazyConnectionFlag, false, "Deprecated: no longer used. Lazy connections are controlled by the server and the NB_LAZY_CONN environment variable.")
|
||||||
_ = upCmd.PersistentFlags().MarkDeprecated(enableLazyConnectionFlag, "no longer used; lazy connections are controlled by the server and the NB_LAZY_CONN environment variable")
|
_ = upCmd.PersistentFlags().MarkDeprecated(enableLazyConnectionFlag, "no longer used; lazy connections are controlled by the server and the NB_LAZY_CONN environment variable")
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|||||||
@@ -41,13 +41,15 @@ func daemonServerOptions(network string) []grpc.ServerOption {
|
|||||||
if network == "tcp" {
|
if network == "tcp" {
|
||||||
log.Warnf("daemon is listening on TCP (%s): callers carry no verifiable identity over TCP, "+
|
log.Warnf("daemon is listening on TCP (%s): callers carry no verifiable identity over TCP, "+
|
||||||
"so privileged operations (SSH root login, SSH auth, enabling the SSH server, management URL changes, "+
|
"so privileged operations (SSH root login, SSH auth, enabling the SSH server, management URL changes, "+
|
||||||
"deregistration) will be denied. Use a unix socket, or npipe:// on Windows", daemonAddr)
|
"deregistration) will be denied, and the SSH JWT cache is neither filled nor served. "+
|
||||||
|
"Use a unix socket, or npipe:// on Windows", daemonAddr)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
creds := ipcauth.NewTransportCredentials() //nolint:staticcheck
|
creds := ipcauth.NewTransportCredentials() //nolint:staticcheck
|
||||||
if creds == nil { //nolint:staticcheck // nil only on platforms without a peer-identity primitive
|
if creds == nil { //nolint:staticcheck // nil only on platforms without a peer-identity primitive
|
||||||
log.Warnf("daemon IPC has no peer-identity primitive on %s: privileged operations will be denied", runtime.GOOS)
|
log.Warnf("daemon IPC has no peer-identity primitive on %s: privileged operations will be denied "+
|
||||||
|
"and the SSH JWT cache is neither filled nor served", runtime.GOOS)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
+172
-63
@@ -2,10 +2,10 @@ package cmd
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os/user"
|
|
||||||
"runtime"
|
"runtime"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -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"
|
||||||
@@ -48,6 +49,8 @@ const (
|
|||||||
profileNameDesc = "profile name to use for the login. If not specified, the last used profile will be used."
|
profileNameDesc = "profile name to use for the login. If not specified, the last used profile will be used."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var errDaemonActiveProfileUnsupported = errors.New("daemon does not support active profile lookup")
|
||||||
|
|
||||||
var (
|
var (
|
||||||
foregroundMode bool
|
foregroundMode bool
|
||||||
dnsLabels []string
|
dnsLabels []string
|
||||||
@@ -122,23 +125,25 @@ func upFunc(cmd *cobra.Command, args []string) error {
|
|||||||
|
|
||||||
pm := profilemanager.NewProfileManager()
|
pm := profilemanager.NewProfileManager()
|
||||||
|
|
||||||
username, err := user.Current()
|
username, err := profilemanager.InvokingUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get current user: %v", err)
|
return fmt.Errorf("get current user: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var activeProf *profilemanager.Profile
|
||||||
var profileSwitched bool
|
var profileSwitched bool
|
||||||
// switch profile if provided
|
// switch profile if provided
|
||||||
if profileName != "" {
|
if profileName != "" {
|
||||||
if err := switchOrCreateProfile(cmd.Context(), pm, profileName, username.Username); err != nil {
|
activeProf, err = switchOrCreateProfile(cmd.Context(), pm, profileName, username.Username)
|
||||||
|
if err != nil {
|
||||||
return fmt.Errorf("switch profile: %v", err)
|
return fmt.Errorf("switch profile: %v", err)
|
||||||
}
|
}
|
||||||
profileSwitched = true
|
profileSwitched = true
|
||||||
}
|
} else {
|
||||||
|
activeProf, err = pm.GetActiveProfile()
|
||||||
activeProf, err := pm.GetActiveProfile()
|
if err != nil {
|
||||||
if err != nil {
|
return fmt.Errorf("get active profile: %v", err)
|
||||||
return fmt.Errorf("get active profile: %v", err)
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if foregroundMode {
|
if foregroundMode {
|
||||||
@@ -150,13 +155,15 @@ func upFunc(cmd *cobra.Command, args []string) error {
|
|||||||
// switchOrCreateProfile switches the active profile to the one identified by
|
// switchOrCreateProfile switches the active profile to the one identified by
|
||||||
// handle, creating it first when it does not exist yet. This restores the
|
// handle, creating it first when it does not exist yet. This restores the
|
||||||
// pre-0.73 behaviour where `netbird up --profile <name>` auto-creates a
|
// pre-0.73 behaviour where `netbird up --profile <name>` auto-creates a
|
||||||
// missing profile instead of failing.
|
// missing profile instead of failing. Returns the daemon-resolved profile so
|
||||||
func switchOrCreateProfile(ctx context.Context, pm *profilemanager.ProfileManager, handle, username string) error {
|
// callers act on it directly instead of re-reading the local state, which is
|
||||||
|
// not updated under sudo.
|
||||||
|
func switchOrCreateProfile(ctx context.Context, pm *profilemanager.ProfileManager, handle, username string) (*profilemanager.Profile, error) {
|
||||||
resolvedID, err := switchProfile(ctx, handle, username)
|
resolvedID, err := switchProfile(ctx, handle, username)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
st, ok := gstatus.FromError(err)
|
st, ok := gstatus.FromError(err)
|
||||||
if !ok || st.Code() != codes.NotFound {
|
if !ok || st.Code() != codes.NotFound {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
// Don't fail immediately on a create error: a concurrent run may
|
// Don't fail immediately on a create error: a concurrent run may
|
||||||
// have created the profile between the NotFound above and this
|
// have created the profile between the NotFound above and this
|
||||||
@@ -165,16 +172,16 @@ func switchOrCreateProfile(ctx context.Context, pm *profilemanager.ProfileManage
|
|||||||
_, createErr := createProfile(ctx, handle, username)
|
_, createErr := createProfile(ctx, handle, username)
|
||||||
if resolvedID, err = switchProfile(ctx, handle, username); err != nil {
|
if resolvedID, err = switchProfile(ctx, handle, username); err != nil {
|
||||||
if createErr != nil {
|
if createErr != nil {
|
||||||
return fmt.Errorf("create profile: %w", createErr)
|
return nil, fmt.Errorf("create profile: %w", createErr)
|
||||||
}
|
}
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := pm.SwitchProfile(resolvedID); err != nil {
|
if err := pm.SwitchProfile(resolvedID); err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
return nil
|
return &profilemanager.Profile{ID: resolvedID}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// createProfile dials the daemon and creates a new profile with the given
|
// createProfile dials the daemon and creates a new profile with the given
|
||||||
@@ -228,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)
|
||||||
|
|
||||||
@@ -302,6 +313,30 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager
|
|||||||
return fmt.Errorf("unable to get daemon status: %v", err)
|
return fmt.Errorf("unable to get daemon status: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Under sudo the invoking user's local active-profile mirror is never
|
||||||
|
// written (the SwitchProfile write is a no-op), and plain root has no
|
||||||
|
// invoking user at all — so the mirror read into activeProf above is stale
|
||||||
|
// or defaulted and must not drive the daemon. With no --profile to make the
|
||||||
|
// choice explicit, take the profile the daemon already holds for this user
|
||||||
|
// instead: it stays on the user's current profile rather than silently
|
||||||
|
// switching to the mirror's default, and refuses when the daemon is on
|
||||||
|
// another user's profile.
|
||||||
|
if profileName == "" && !profilemanager.MirrorIsAuthoritative() {
|
||||||
|
u, err := profilemanager.InvokingUser()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("get current user: %v", err)
|
||||||
|
}
|
||||||
|
resolved, err := daemonActiveProfileForUser(ctx, client, u.Username)
|
||||||
|
switch {
|
||||||
|
case errors.Is(err, errDaemonActiveProfileUnsupported):
|
||||||
|
log.Warnf("keeping the locally resolved profile: %v", err)
|
||||||
|
case err != nil:
|
||||||
|
return err
|
||||||
|
default:
|
||||||
|
activeProf = resolved
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if status.Status == string(internal.StatusConnected) {
|
if status.Status == string(internal.StatusConnected) {
|
||||||
if !profileSwitched {
|
if !profileSwitched {
|
||||||
cmd.Println("Already connected")
|
cmd.Println("Already connected")
|
||||||
@@ -314,7 +349,7 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
username, err := user.Current()
|
username, err := profilemanager.InvokingUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("get current user: %v", err)
|
return fmt.Errorf("get current user: %v", err)
|
||||||
}
|
}
|
||||||
@@ -398,26 +433,21 @@ func doDaemonUp(ctx context.Context, cmd *cobra.Command, client proto.DaemonServ
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, profileName, username string) *proto.SetConfigRequest {
|
// setBoolPtrIfChanged points dst at a copy of val when the named bool flag was
|
||||||
var req proto.SetConfigRequest
|
// explicitly set on cmd. It collapses the repeated
|
||||||
req.ProfileName = profileName
|
// "if cmd.Flag(x).Changed { field = &val }" pattern in the request builders into
|
||||||
req.Username = username
|
// a single call, keeping their cognitive complexity within bounds.
|
||||||
|
func setBoolPtrIfChanged(cmd *cobra.Command, name string, dst **bool, val bool) {
|
||||||
req.ManagementUrl = managementURL
|
if cmd.Flag(name).Changed {
|
||||||
req.AdminURL = adminURL
|
dst2 := val
|
||||||
req.NatExternalIPs = natExternalIPs
|
*dst = &dst2
|
||||||
req.CustomDNSAddress = customDNSAddressConverted
|
|
||||||
req.ExtraIFaceBlacklist = extraIFaceBlackList
|
|
||||||
req.DnsLabels = dnsLabelsValidated.ToPunycodeList()
|
|
||||||
req.CleanDNSLabels = dnsLabels != nil && len(dnsLabels) == 0
|
|
||||||
req.CleanNATExternalIPs = natExternalIPs != nil && len(natExternalIPs) == 0
|
|
||||||
|
|
||||||
if cmd.Flag(enableRosenpassFlag).Changed {
|
|
||||||
req.RosenpassEnabled = &rosenpassEnabled
|
|
||||||
}
|
|
||||||
if cmd.Flag(rosenpassPermissiveFlag).Changed {
|
|
||||||
req.RosenpassPermissive = &rosenpassPermissive
|
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// setSSHSetConfigFields copies the SSH server flags the user actually
|
||||||
|
// passed into req, leaving the rest unset so the daemon keeps the
|
||||||
|
// persisted values.
|
||||||
|
func setSSHSetConfigFields(req *proto.SetConfigRequest, cmd *cobra.Command) {
|
||||||
if cmd.Flag(serverSSHAllowedFlag).Changed {
|
if cmd.Flag(serverSSHAllowedFlag).Changed {
|
||||||
req.ServerSSHAllowed = &serverSSHAllowed
|
req.ServerSSHAllowed = &serverSSHAllowed
|
||||||
}
|
}
|
||||||
@@ -440,6 +470,31 @@ func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, pro
|
|||||||
sshJWTCacheTTL32 := int32(sshJWTCacheTTL)
|
sshJWTCacheTTL32 := int32(sshJWTCacheTTL)
|
||||||
req.SshJWTCacheTTL = &sshJWTCacheTTL32
|
req.SshJWTCacheTTL = &sshJWTCacheTTL32
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, profileName, username string) *proto.SetConfigRequest {
|
||||||
|
var req proto.SetConfigRequest
|
||||||
|
req.ProfileName = profileName
|
||||||
|
req.Username = username
|
||||||
|
|
||||||
|
req.ManagementUrl = managementURL
|
||||||
|
req.AdminURL = adminURL
|
||||||
|
req.NatExternalIPs = natExternalIPs
|
||||||
|
req.CustomDNSAddress = customDNSAddressConverted
|
||||||
|
req.ExtraIFaceBlacklist = extraIFaceBlackList
|
||||||
|
req.DnsLabels = dnsLabelsValidated.ToPunycodeList()
|
||||||
|
req.CleanDNSLabels = dnsLabels != nil && len(dnsLabels) == 0
|
||||||
|
req.CleanNATExternalIPs = natExternalIPs != nil && len(natExternalIPs) == 0
|
||||||
|
|
||||||
|
if cmd.Flag(enableRosenpassFlag).Changed {
|
||||||
|
req.RosenpassEnabled = &rosenpassEnabled
|
||||||
|
}
|
||||||
|
if cmd.Flag(rosenpassPermissiveFlag).Changed {
|
||||||
|
req.RosenpassPermissive = &rosenpassPermissive
|
||||||
|
}
|
||||||
|
setSSHSetConfigFields(&req, cmd)
|
||||||
|
setBoolPtrIfChanged(cmd, remoteJobsAllowedFlag, &req.RemoteJobsAllowed, remoteJobsAllowed)
|
||||||
|
|
||||||
if cmd.Flag(interfaceNameFlag).Changed {
|
if cmd.Flag(interfaceNameFlag).Changed {
|
||||||
if err := parseInterfaceName(interfaceName); err != nil {
|
if err := parseInterfaceName(interfaceName); err != nil {
|
||||||
log.Errorf("parse interface name: %v", err)
|
log.Errorf("parse interface name: %v", err)
|
||||||
@@ -499,6 +554,13 @@ func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, pro
|
|||||||
req.DisableIpv6 = &disableIPv6
|
req.DisableIpv6 = &disableIPv6
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if cmd.Flag(enableLocalMetricsFlag).Changed {
|
||||||
|
req.EnableLocalMetrics = &localMetricsEnabled
|
||||||
|
}
|
||||||
|
if cmd.Flag(localMetricsAddressFlag).Changed {
|
||||||
|
req.LocalMetricsAddress = &localMetricsAddr
|
||||||
|
}
|
||||||
|
|
||||||
return &req
|
return &req
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -523,6 +585,7 @@ func setupConfig(customDNSAddressConverted []byte, cmd *cobra.Command, configFil
|
|||||||
if cmd.Flag(serverSSHAllowedFlag).Changed {
|
if cmd.Flag(serverSSHAllowedFlag).Changed {
|
||||||
ic.ServerSSHAllowed = &serverSSHAllowed
|
ic.ServerSSHAllowed = &serverSSHAllowed
|
||||||
}
|
}
|
||||||
|
setBoolPtrIfChanged(cmd, remoteJobsAllowedFlag, &ic.RemoteJobsAllowed, remoteJobsAllowed)
|
||||||
|
|
||||||
if cmd.Flag(enableSSHRootFlag).Changed {
|
if cmd.Flag(enableSSHRootFlag).Changed {
|
||||||
ic.EnableSSHRoot = &enableSSHRoot
|
ic.EnableSSHRoot = &enableSSHRoot
|
||||||
@@ -616,9 +679,45 @@ func setupConfig(customDNSAddressConverted []byte, cmd *cobra.Command, configFil
|
|||||||
ic.DisableIPv6 = &disableIPv6
|
ic.DisableIPv6 = &disableIPv6
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if cmd.Flag(enableLocalMetricsFlag).Changed {
|
||||||
|
ic.LocalMetricsEnabled = &localMetricsEnabled
|
||||||
|
}
|
||||||
|
|
||||||
|
if cmd.Flag(localMetricsAddressFlag).Changed {
|
||||||
|
ic.LocalMetricsAddress = &localMetricsAddr
|
||||||
|
}
|
||||||
|
|
||||||
return &ic, nil
|
return &ic, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// setSSHLoginFields copies the SSH server flags the user actually passed
|
||||||
|
// into req, leaving the rest unset so the daemon keeps the persisted
|
||||||
|
// values.
|
||||||
|
func setSSHLoginFields(req *proto.LoginRequest, cmd *cobra.Command) {
|
||||||
|
if cmd.Flag(serverSSHAllowedFlag).Changed {
|
||||||
|
req.ServerSSHAllowed = &serverSSHAllowed
|
||||||
|
}
|
||||||
|
if cmd.Flag(enableSSHRootFlag).Changed {
|
||||||
|
req.EnableSSHRoot = &enableSSHRoot
|
||||||
|
}
|
||||||
|
if cmd.Flag(enableSSHSFTPFlag).Changed {
|
||||||
|
req.EnableSSHSFTP = &enableSSHSFTP
|
||||||
|
}
|
||||||
|
if cmd.Flag(enableSSHLocalPortForwardFlag).Changed {
|
||||||
|
req.EnableSSHLocalPortForwarding = &enableSSHLocalPortForward
|
||||||
|
}
|
||||||
|
if cmd.Flag(enableSSHRemotePortForwardFlag).Changed {
|
||||||
|
req.EnableSSHRemotePortForwarding = &enableSSHRemotePortForward
|
||||||
|
}
|
||||||
|
if cmd.Flag(disableSSHAuthFlag).Changed {
|
||||||
|
req.DisableSSHAuth = &disableSSHAuth
|
||||||
|
}
|
||||||
|
if cmd.Flag(sshJWTCacheTTLFlag).Changed {
|
||||||
|
sshJWTCacheTTL32 := int32(sshJWTCacheTTL)
|
||||||
|
req.SshJWTCacheTTL = &sshJWTCacheTTL32
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func setupLoginRequest(providedSetupKey string, customDNSAddressConverted []byte, cmd *cobra.Command) (*proto.LoginRequest, error) {
|
func setupLoginRequest(providedSetupKey string, customDNSAddressConverted []byte, cmd *cobra.Command) (*proto.LoginRequest, error) {
|
||||||
loginRequest := proto.LoginRequest{
|
loginRequest := proto.LoginRequest{
|
||||||
SetupKey: providedSetupKey,
|
SetupKey: providedSetupKey,
|
||||||
@@ -645,39 +744,21 @@ func setupLoginRequest(providedSetupKey string, customDNSAddressConverted []byte
|
|||||||
loginRequest.RosenpassPermissive = &rosenpassPermissive
|
loginRequest.RosenpassPermissive = &rosenpassPermissive
|
||||||
}
|
}
|
||||||
|
|
||||||
if cmd.Flag(serverSSHAllowedFlag).Changed {
|
setSSHLoginFields(&loginRequest, cmd)
|
||||||
loginRequest.ServerSSHAllowed = &serverSSHAllowed
|
setBoolPtrIfChanged(cmd, remoteJobsAllowedFlag, &loginRequest.RemoteJobsAllowed, remoteJobsAllowed)
|
||||||
}
|
|
||||||
|
|
||||||
if cmd.Flag(enableSSHRootFlag).Changed {
|
|
||||||
loginRequest.EnableSSHRoot = &enableSSHRoot
|
|
||||||
}
|
|
||||||
|
|
||||||
if cmd.Flag(enableSSHSFTPFlag).Changed {
|
|
||||||
loginRequest.EnableSSHSFTP = &enableSSHSFTP
|
|
||||||
}
|
|
||||||
|
|
||||||
if cmd.Flag(enableSSHLocalPortForwardFlag).Changed {
|
|
||||||
loginRequest.EnableSSHLocalPortForwarding = &enableSSHLocalPortForward
|
|
||||||
}
|
|
||||||
|
|
||||||
if cmd.Flag(enableSSHRemotePortForwardFlag).Changed {
|
|
||||||
loginRequest.EnableSSHRemotePortForwarding = &enableSSHRemotePortForward
|
|
||||||
}
|
|
||||||
|
|
||||||
if cmd.Flag(disableSSHAuthFlag).Changed {
|
|
||||||
loginRequest.DisableSSHAuth = &disableSSHAuth
|
|
||||||
}
|
|
||||||
|
|
||||||
if cmd.Flag(sshJWTCacheTTLFlag).Changed {
|
|
||||||
sshJWTCacheTTL32 := int32(sshJWTCacheTTL)
|
|
||||||
loginRequest.SshJWTCacheTTL = &sshJWTCacheTTL32
|
|
||||||
}
|
|
||||||
|
|
||||||
if cmd.Flag(disableAutoConnectFlag).Changed {
|
if cmd.Flag(disableAutoConnectFlag).Changed {
|
||||||
loginRequest.DisableAutoConnect = &autoConnectDisabled
|
loginRequest.DisableAutoConnect = &autoConnectDisabled
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if cmd.Flag(enableLocalMetricsFlag).Changed {
|
||||||
|
loginRequest.EnableLocalMetrics = &localMetricsEnabled
|
||||||
|
}
|
||||||
|
|
||||||
|
if cmd.Flag(localMetricsAddressFlag).Changed {
|
||||||
|
loginRequest.LocalMetricsAddress = &localMetricsAddr
|
||||||
|
}
|
||||||
|
|
||||||
if cmd.Flag(interfaceNameFlag).Changed {
|
if cmd.Flag(interfaceNameFlag).Changed {
|
||||||
if err := parseInterfaceName(interfaceName); err != nil {
|
if err := parseInterfaceName(interfaceName); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -849,3 +930,31 @@ func isValidAddrPort(input string) bool {
|
|||||||
_, err := netip.ParseAddrPort(input)
|
_, err := netip.ParseAddrPort(input)
|
||||||
return err == nil
|
return err == nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// daemonActiveProfileForUser returns the profile the daemon currently holds for
|
||||||
|
// username, for the no --profile case where the local mirror is not
|
||||||
|
// authoritative (sudo or plain root). It returns that profile when the daemon
|
||||||
|
// owns it for this user or when the profile is unowned (empty username, as on a
|
||||||
|
// fresh install), so the caller acts on the daemon's real state instead of the
|
||||||
|
// stale mirror. It denies with a --profile hint when the daemon is on another
|
||||||
|
// user's profile, when the lookup fails, or when the daemon reports no active
|
||||||
|
// profile. Returns errDaemonActiveProfileUnsupported when the daemon predates
|
||||||
|
// the RPC; the caller keeps the mirror-derived profile in that case.
|
||||||
|
func daemonActiveProfileForUser(ctx context.Context, client proto.DaemonServiceClient, username string) (*profilemanager.Profile, error) {
|
||||||
|
active, err := client.GetActiveProfile(ctx, &proto.GetActiveProfileRequest{})
|
||||||
|
if err != nil {
|
||||||
|
if st, ok := gstatus.FromError(err); ok && st.Code() == codes.Unimplemented {
|
||||||
|
return nil, fmt.Errorf("%w: %v", errDaemonActiveProfileUnsupported, err)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("pass --profile to choose the profile explicitly: the daemon's active profile could not be verified: %v", err)
|
||||||
|
}
|
||||||
|
if active.GetId() == "" {
|
||||||
|
return nil, fmt.Errorf("pass --profile to choose the profile explicitly: the daemon reported no active profile")
|
||||||
|
}
|
||||||
|
if active.GetUsername() != "" && active.GetUsername() != username {
|
||||||
|
return nil, fmt.Errorf(
|
||||||
|
"pass --profile to choose the profile explicitly: the daemon's active profile is %q (user %q) but this invocation runs for %q",
|
||||||
|
active.GetProfileName(), active.GetUsername(), username)
|
||||||
|
}
|
||||||
|
return &profilemanager.Profile{ID: profilemanager.ID(active.GetId())}, nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,88 @@
|
|||||||
|
package cmd
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"google.golang.org/grpc"
|
||||||
|
"google.golang.org/grpc/codes"
|
||||||
|
gstatus "google.golang.org/grpc/status"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
"github.com/netbirdio/netbird/client/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
type fakeActiveProfileClient struct {
|
||||||
|
proto.DaemonServiceClient
|
||||||
|
resp *proto.GetActiveProfileResponse
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeActiveProfileClient) GetActiveProfile(_ context.Context, _ *proto.GetActiveProfileRequest, _ ...grpc.CallOption) (*proto.GetActiveProfileResponse, error) {
|
||||||
|
return f.resp, f.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDaemonActiveProfileForUserReturnsOwnProfile(t *testing.T) {
|
||||||
|
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{Id: "default", ProfileName: "default", Username: "root"}}
|
||||||
|
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, prof)
|
||||||
|
assert.Equal(t, profilemanager.ID("default"), prof.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDaemonActiveProfileForUserReturnsUnownedProfile(t *testing.T) {
|
||||||
|
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{Id: "default", ProfileName: "default", Username: ""}}
|
||||||
|
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, prof)
|
||||||
|
assert.Equal(t, profilemanager.ID("default"), prof.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDaemonActiveProfileForUserKeepsDaemonProfileOverStaleMirror(t *testing.T) {
|
||||||
|
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{Id: "ab12", ProfileName: "work", Username: "misha"}}
|
||||||
|
prof, err := daemonActiveProfileForUser(context.Background(), client, "misha")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, prof)
|
||||||
|
assert.Equal(t, profilemanager.ID("ab12"), prof.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDaemonActiveProfileForUserRejectsOtherUsersProfile(t *testing.T) {
|
||||||
|
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{Id: "ab12", ProfileName: "work", Username: "misha"}}
|
||||||
|
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Nil(t, prof)
|
||||||
|
assert.Contains(t, err.Error(), "--profile")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDaemonActiveProfileForUserRejectsOtherUsersDefaultProfile(t *testing.T) {
|
||||||
|
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{Id: "default", ProfileName: "default", Username: "misha"}}
|
||||||
|
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Nil(t, prof)
|
||||||
|
assert.Contains(t, err.Error(), "--profile")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDaemonActiveProfileForUserRejectsLookupError(t *testing.T) {
|
||||||
|
client := &fakeActiveProfileClient{err: gstatus.Error(codes.Internal, "boom")}
|
||||||
|
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Nil(t, prof)
|
||||||
|
assert.Contains(t, err.Error(), "--profile")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDaemonActiveProfileForUserRejectsEmptyResponse(t *testing.T) {
|
||||||
|
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{}}
|
||||||
|
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Nil(t, prof)
|
||||||
|
assert.Contains(t, err.Error(), "--profile")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDaemonActiveProfileForUserKeepsMirrorWhenDaemonWithoutRPC(t *testing.T) {
|
||||||
|
client := &fakeActiveProfileClient{err: gstatus.Error(codes.Unimplemented, "unknown method")}
|
||||||
|
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
|
||||||
|
require.ErrorIs(t, err, errDaemonActiveProfileUnsupported)
|
||||||
|
assert.Nil(t, prof)
|
||||||
|
}
|
||||||
@@ -0,0 +1,77 @@
|
|||||||
|
#!/bin/sh
|
||||||
|
set -eu
|
||||||
|
|
||||||
|
if [ "$#" -lt 2 ]; then
|
||||||
|
printf '%s\n' "usage: $0 OUTPUT_DIRECTORY GOARCH..." >&2
|
||||||
|
exit 2
|
||||||
|
fi
|
||||||
|
|
||||||
|
repo_root=$(CDPATH= cd -- "$(dirname "$0")/.." && pwd)
|
||||||
|
output_name=$(basename "$1")
|
||||||
|
if [ -z "$output_name" ] || [ "$output_name" = "." ] ||
|
||||||
|
[ "$output_name" = ".." ] || [ "$output_name" = "/" ]; then
|
||||||
|
printf '%s\n' "OUTPUT_DIRECTORY must name a directory" >&2
|
||||||
|
exit 2
|
||||||
|
fi
|
||||||
|
output_parent=$(CDPATH= cd -- "$(dirname "$1")" && pwd)
|
||||||
|
output="$output_parent/$output_name"
|
||||||
|
shift
|
||||||
|
modules=$(mktemp "${TMPDIR:-/tmp}/netbird-client-licenses.modules.XXXXXX")
|
||||||
|
sorted_modules=$(mktemp "${TMPDIR:-/tmp}/netbird-client-licenses.sorted.XXXXXX")
|
||||||
|
trap 'rm -f "$modules" "$sorted_modules"' EXIT HUP INT TERM
|
||||||
|
|
||||||
|
if [ -e "$output" ] || [ -L "$output" ]; then
|
||||||
|
printf 'output directory already exists: %s\n' "$output" >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
mkdir "$output"
|
||||||
|
mkdir "$output/third_party"
|
||||||
|
|
||||||
|
cp "$repo_root/LICENSE" "$output/BSD-3-Clause.txt"
|
||||||
|
|
||||||
|
cd "$repo_root"
|
||||||
|
for arch in "$@"; do
|
||||||
|
GOOS=${GOOS:-linux} GOARCH="$arch" CGO_ENABLED=${CGO_ENABLED:-0} \
|
||||||
|
go list -deps -f '{{with .Module}}{{if .Replace}}{{.Replace.Path}}{{"\t"}}{{.Replace.Version}}{{"\t"}}{{.Replace.Dir}}{{else}}{{.Path}}{{"\t"}}{{.Version}}{{"\t"}}{{.Dir}}{{end}}{{end}}' -tags load_wgnt_from_rsrc ./client >>"$modules"
|
||||||
|
done
|
||||||
|
LC_ALL=C sort -u "$modules" >"$sorted_modules"
|
||||||
|
|
||||||
|
goroot=$(go env GOROOT)
|
||||||
|
for term in LICENSE PATENTS; do
|
||||||
|
if [ ! -f "$goroot/$term" ]; then
|
||||||
|
printf 'missing Go standard-library term: %s\n' "$goroot/$term" >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
cp "$goroot/$term" "$output/Go-$term"
|
||||||
|
done
|
||||||
|
|
||||||
|
while IFS=' ' read -r module version module_dir; do
|
||||||
|
[ -n "$module" ] || continue
|
||||||
|
[ "$module" = "github.com/netbirdio/netbird" ] && continue
|
||||||
|
|
||||||
|
if [ -z "$version" ] || [ ! -d "$module_dir" ]; then
|
||||||
|
printf 'cannot collect terms for module %s at version %s\n' "$module" "$version" >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
|
destination="$output/third_party/$module/$version"
|
||||||
|
mkdir -p "$destination"
|
||||||
|
printf 'module: %s\nversion: %s\n' "$module" "$version" >"$destination/MODULE"
|
||||||
|
|
||||||
|
found=false
|
||||||
|
for term in \
|
||||||
|
"$module_dir"/LICENSE* "$module_dir"/License* "$module_dir"/license* \
|
||||||
|
"$module_dir"/LICENCE* "$module_dir"/Licence* "$module_dir"/licence* \
|
||||||
|
"$module_dir"/COPYING* "$module_dir"/Copying* "$module_dir"/copying* \
|
||||||
|
"$module_dir"/NOTICE* "$module_dir"/Notice* "$module_dir"/notice* \
|
||||||
|
"$module_dir"/PATENTS* "$module_dir"/Patents* "$module_dir"/patents*; do
|
||||||
|
[ -f "$term" ] || continue
|
||||||
|
cp "$term" "$destination/"
|
||||||
|
found=true
|
||||||
|
done
|
||||||
|
|
||||||
|
if [ "$found" = false ]; then
|
||||||
|
printf 'no root license terms found for module %s at %s\n' "$module" "$module_dir" >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
done <"$sorted_modules"
|
||||||
@@ -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
|
||||||
|
|||||||
@@ -0,0 +1,11 @@
|
|||||||
|
//go:build android || (!linux && !windows)
|
||||||
|
|
||||||
|
package firewall
|
||||||
|
|
||||||
|
import "github.com/netbirdio/netbird/client/firewall/uspfilter"
|
||||||
|
|
||||||
|
// interfaceAllower returns no allower: these platforms have no host firewall to
|
||||||
|
// open for the interface.
|
||||||
|
func interfaceAllower(IFaceMapper, uint16) uspfilter.InterfaceAllower {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package firewall
|
||||||
|
|
||||||
|
import "github.com/netbirdio/netbird/client/firewall/uspfilter"
|
||||||
|
|
||||||
|
// interfaceAllower returns the Windows netsh-based interface allower.
|
||||||
|
func interfaceAllower(iface IFaceMapper, _ uint16) uspfilter.InterfaceAllower {
|
||||||
|
return uspfilter.NewWindowsInterfaceAllower(iface)
|
||||||
|
}
|
||||||
@@ -6,8 +6,6 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"runtime"
|
"runtime"
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
|
||||||
|
|
||||||
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
"github.com/netbirdio/netbird/client/firewall/uspfilter"
|
"github.com/netbirdio/netbird/client/firewall/uspfilter"
|
||||||
nftypes "github.com/netbirdio/netbird/client/internal/netflow/types"
|
nftypes "github.com/netbirdio/netbird/client/internal/netflow/types"
|
||||||
@@ -21,13 +19,11 @@ func NewFirewall(iface IFaceMapper, _ *statemanager.Manager, flowLogger nftypes.
|
|||||||
}
|
}
|
||||||
|
|
||||||
// use userspace packet filtering firewall
|
// use userspace packet filtering firewall
|
||||||
fm, err := uspfilter.Create(iface, disableServerRoutes, flowLogger, mtu)
|
return uspfilter.Create(uspfilter.Config{
|
||||||
if err != nil {
|
IFace: iface,
|
||||||
return nil, err
|
DisableServerRoutes: disableServerRoutes,
|
||||||
}
|
FlowLogger: flowLogger,
|
||||||
err = fm.AllowNetbird()
|
MTU: mtu,
|
||||||
if err != nil {
|
InterfaceAllower: interfaceAllower(iface, mtu),
|
||||||
log.Warnf("failed to allow netbird interface traffic: %v", err)
|
})
|
||||||
}
|
|
||||||
return fm, nil
|
|
||||||
}
|
}
|
||||||
|
|||||||
+126
-75
@@ -16,6 +16,7 @@ import (
|
|||||||
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
nbnftables "github.com/netbirdio/netbird/client/firewall/nftables"
|
nbnftables "github.com/netbirdio/netbird/client/firewall/nftables"
|
||||||
"github.com/netbirdio/netbird/client/firewall/uspfilter"
|
"github.com/netbirdio/netbird/client/firewall/uspfilter"
|
||||||
|
"github.com/netbirdio/netbird/client/iface/netstack"
|
||||||
nftypes "github.com/netbirdio/netbird/client/internal/netflow/types"
|
nftypes "github.com/netbirdio/netbird/client/internal/netflow/types"
|
||||||
"github.com/netbirdio/netbird/client/internal/statemanager"
|
"github.com/netbirdio/netbird/client/internal/statemanager"
|
||||||
)
|
)
|
||||||
@@ -29,47 +30,107 @@ const (
|
|||||||
NFTABLES
|
NFTABLES
|
||||||
)
|
)
|
||||||
|
|
||||||
// SKIP_NFTABLES_ENV is the environment variable to skip nftables check
|
// SkipNftablesEnv is the environment variable to skip nftables check
|
||||||
const SKIP_NFTABLES_ENV = "NB_SKIP_NFTABLES_CHECK"
|
const SkipNftablesEnv = "NB_SKIP_NFTABLES_CHECK"
|
||||||
|
|
||||||
|
// errNoFirewallManager indicates no kernel firewall backend is present,
|
||||||
|
// as opposed to a backend that exists but failed to create or initialize.
|
||||||
|
var errNoFirewallManager = errors.New("no firewall manager found")
|
||||||
|
|
||||||
// FWType is the type for the firewall type
|
// FWType is the type for the firewall type
|
||||||
type FWType int
|
type FWType int
|
||||||
|
|
||||||
func NewFirewall(iface IFaceMapper, stateManager *statemanager.Manager, flowLogger nftypes.FlowLogger, disableServerRoutes bool, mtu uint16) (firewall.Manager, error) {
|
func NewFirewall(iface IFaceMapper, stateManager *statemanager.Manager, flowLogger nftypes.FlowLogger, disableServerRoutes bool, mtu uint16) (firewall.Manager, error) {
|
||||||
// We run in userspace mode and force userspace firewall was requested. We don't attempt native firewall.
|
// Userspace firewall without a native counterpart: routing is handled
|
||||||
if iface.IsUserspaceBind() && forceUserspaceFirewall() {
|
// entirely in userspace. The interface is opened in the kernel's foreign
|
||||||
log.Info("forcing userspace firewall")
|
// filter chains via a table-less allower, except in netstack mode where no
|
||||||
return createUserspaceFirewall(iface, nil, disableServerRoutes, flowLogger, mtu)
|
// kernel interface exists.
|
||||||
|
if netstack.IsEnabled() || (iface.IsUserspaceBind() && forceUserspaceFirewall()) {
|
||||||
|
if netstack.IsEnabled() {
|
||||||
|
log.Info("netstack mode, using userspace firewall")
|
||||||
|
} else {
|
||||||
|
log.Info("forcing userspace firewall")
|
||||||
|
}
|
||||||
|
cfg := uspfilter.Config{
|
||||||
|
IFace: iface,
|
||||||
|
DisableServerRoutes: disableServerRoutes,
|
||||||
|
FlowLogger: flowLogger,
|
||||||
|
MTU: mtu,
|
||||||
|
InterfaceAllower: interfaceAllower(iface, mtu),
|
||||||
|
}
|
||||||
|
|
||||||
|
return uspfilter.Create(cfg)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Use native firewall for either kernel or userspace, the interface appears identical to netfilter
|
// Use native firewall for either kernel or userspace, the interface appears identical to netfilter
|
||||||
fm, err := createNativeFirewall(iface, stateManager, disableServerRoutes, mtu)
|
fm, err := createNativeFirewall(iface, stateManager, mtu)
|
||||||
|
switch {
|
||||||
// Kernel cannot fall back to anything else, need to return error
|
case err == nil && !iface.IsUserspaceBind():
|
||||||
if !iface.IsUserspaceBind() {
|
// Nothing to do, fall through
|
||||||
return fm, err
|
case err == nil && iface.IsUserspaceBind():
|
||||||
}
|
// Native firewall handles packet filtering, but the userspace WireGuard bind
|
||||||
|
// needs a device filter for DNS interception hooks. Install a minimal
|
||||||
// Fall back to the userspace packet filter if native is unavailable
|
// hooks-only filter that passes all traffic through to the kernel firewall.
|
||||||
if err != nil {
|
if err := iface.SetFilter(&uspfilter.HooksFilter{}); err != nil {
|
||||||
log.Warnf("failed to create native firewall: %v. Proceeding with userspace", err)
|
log.Warnf("failed to set hooks filter, DNS via memory hooks will not work: %v", err)
|
||||||
return createUserspaceFirewall(iface, nil, disableServerRoutes, flowLogger, mtu)
|
}
|
||||||
}
|
case err != nil && !iface.IsUserspaceBind():
|
||||||
|
// Kernel cannot fall back to anything else, need to return error
|
||||||
// Native firewall handles packet filtering, but the userspace WireGuard bind
|
return nil, err
|
||||||
// needs a device filter for DNS interception hooks. Install a minimal
|
case err != nil && iface.IsUserspaceBind():
|
||||||
// hooks-only filter that passes all traffic through to the kernel firewall.
|
// Fall back to the userspace packet filter if native is unavailable
|
||||||
if err := iface.SetFilter(&uspfilter.HooksFilter{}); err != nil {
|
logNativeFirewallUnavailable(err)
|
||||||
log.Warnf("failed to set hooks filter, DNS via memory hooks will not work: %v", err)
|
return uspfilter.Create(uspfilter.Config{
|
||||||
|
IFace: iface,
|
||||||
|
DisableServerRoutes: disableServerRoutes,
|
||||||
|
FlowLogger: flowLogger,
|
||||||
|
MTU: mtu,
|
||||||
|
InterfaceAllower: interfaceAllower(iface, mtu),
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
return fm, nil
|
return fm, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func createNativeFirewall(iface IFaceMapper, stateManager *statemanager.Manager, routes bool, mtu uint16) (firewall.Manager, error) {
|
// interfaceAllower selects how the userspace firewall opens the interface in
|
||||||
|
// foreign kernel chains: nftables when available (which also opens foreign nft
|
||||||
|
// tables), else iptables (the legacy fallback, filter INPUT only), else nil.
|
||||||
|
// firewalld trust is applied separately by the manager. Netstack has no kernel
|
||||||
|
// interface to open.
|
||||||
|
func interfaceAllower(iface IFaceMapper, mtu uint16) uspfilter.InterfaceAllower {
|
||||||
|
if netstack.IsEnabled() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
nftAllower, err := nbnftables.NewInterfaceAllower(iface, mtu)
|
||||||
|
if err == nil {
|
||||||
|
return nftAllower
|
||||||
|
}
|
||||||
|
log.Infof("no nftables interface allower: %v", err)
|
||||||
|
|
||||||
|
iptAllower, err := nbiptables.NewInterfaceAllower(iface)
|
||||||
|
if err == nil {
|
||||||
|
return iptAllower
|
||||||
|
}
|
||||||
|
log.Infof("no iptables interface allower: %v", err)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// logNativeFirewallUnavailable logs the fallback to userspace at info level
|
||||||
|
// when no kernel firewall backend exists, and at warn level otherwise.
|
||||||
|
func logNativeFirewallUnavailable(err error) {
|
||||||
|
if errors.Is(err, errNoFirewallManager) {
|
||||||
|
log.Infof("no native firewall backend available: %v. Proceeding with userspace", err)
|
||||||
|
} else {
|
||||||
|
log.Warnf("failed to create native firewall: %v. Proceeding with userspace", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func createNativeFirewall(iface IFaceMapper, stateManager *statemanager.Manager, mtu uint16) (firewall.Manager, error) {
|
||||||
fm, err := createFW(iface, mtu)
|
fm, err := createFW(iface, mtu)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("create firewall: %s", err)
|
return nil, fmt.Errorf("create firewall: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err = fm.Init(stateManager); err != nil {
|
if err = fm.Init(stateManager); err != nil {
|
||||||
@@ -88,29 +149,10 @@ func createFW(iface IFaceMapper, mtu uint16) (firewall.Manager, error) {
|
|||||||
log.Info("creating an nftables firewall manager")
|
log.Info("creating an nftables firewall manager")
|
||||||
return nbnftables.Create(iface, mtu)
|
return nbnftables.Create(iface, mtu)
|
||||||
default:
|
default:
|
||||||
log.Info("no firewall manager found, trying to use userspace packet filtering firewall")
|
return nil, errNoFirewallManager
|
||||||
return nil, errors.New("no firewall manager found")
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func createUserspaceFirewall(iface IFaceMapper, fm firewall.Manager, disableServerRoutes bool, flowLogger nftypes.FlowLogger, mtu uint16) (firewall.Manager, error) {
|
|
||||||
var errUsp error
|
|
||||||
if fm != nil {
|
|
||||||
fm, errUsp = uspfilter.CreateWithNativeFirewall(iface, fm, disableServerRoutes, flowLogger, mtu)
|
|
||||||
} else {
|
|
||||||
fm, errUsp = uspfilter.Create(iface, disableServerRoutes, flowLogger, mtu)
|
|
||||||
}
|
|
||||||
|
|
||||||
if errUsp != nil {
|
|
||||||
return nil, fmt.Errorf("create userspace firewall: %s", errUsp)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := fm.AllowNetbird(); err != nil {
|
|
||||||
log.Errorf("failed to allow netbird interface traffic: %v", err)
|
|
||||||
}
|
|
||||||
return fm, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// check returns the firewall type based on common lib checks. It returns UNKNOWN if no firewall is found.
|
// check returns the firewall type based on common lib checks. It returns UNKNOWN if no firewall is found.
|
||||||
func check() FWType {
|
func check() FWType {
|
||||||
useIPTABLES := false
|
useIPTABLES := false
|
||||||
@@ -132,35 +174,38 @@ func check() FWType {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
nf := nftables.Conn{}
|
// Honor the skip env before probing nftables at all.
|
||||||
if chains, err := nf.ListChains(); err == nil && os.Getenv(SKIP_NFTABLES_ENV) != "true" {
|
if os.Getenv(SkipNftablesEnv) != "true" {
|
||||||
if !useIPTABLES {
|
nf := nftables.Conn{}
|
||||||
return NFTABLES
|
if chains, err := nf.ListChains(); err == nil {
|
||||||
}
|
if !useIPTABLES {
|
||||||
|
|
||||||
// search for chains where table is filter
|
|
||||||
// if we find one, we assume that nftables manager can be used with iptables
|
|
||||||
for _, chain := range chains {
|
|
||||||
if chain.Table.Name == "filter" {
|
|
||||||
return NFTABLES
|
return NFTABLES
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
// check tables for the following constraints:
|
// search for chains where table is filter
|
||||||
// 1. there is no chain in nftables for the filter table and there is at least one chain in iptables, we assume that nftables manager can not be used
|
// if we find one, we assume that nftables manager can be used with iptables
|
||||||
// 2. there is no tables or more than one table, we assume that nftables manager can be used
|
for _, chain := range chains {
|
||||||
// 3. there is only one table and its name is filter, we assume that nftables manager can not be used, since there was no chain in it
|
if chain.Table.Name == "filter" {
|
||||||
// 4. if we find an error we log and continue with iptables check
|
return NFTABLES
|
||||||
nbTablesList, err := nf.ListTables()
|
}
|
||||||
switch {
|
}
|
||||||
case err == nil && len(iptablesChains) > 0:
|
|
||||||
return IPTABLES
|
// check tables for the following constraints:
|
||||||
case err == nil && len(nbTablesList) != 1:
|
// 1. there is no chain in nftables for the filter table and there is at least one chain in iptables, we assume that nftables manager can not be used
|
||||||
return NFTABLES
|
// 2. there is no tables or more than one table, we assume that nftables manager can be used
|
||||||
case err == nil && len(nbTablesList) == 1 && nbTablesList[0].Name == "filter":
|
// 3. there is only one table and its name is filter, we assume that nftables manager can not be used, since there was no chain in it
|
||||||
return IPTABLES
|
// 4. if we find an error we log and continue with iptables check
|
||||||
case err != nil:
|
nbTablesList, err := nf.ListTables()
|
||||||
log.Errorf("failed to list nftables tables on fw manager discovery: %s", err)
|
switch {
|
||||||
|
case err == nil && len(iptablesChains) > 0:
|
||||||
|
return IPTABLES
|
||||||
|
case err == nil && len(nbTablesList) != 1:
|
||||||
|
return NFTABLES
|
||||||
|
case err == nil && len(nbTablesList) == 1 && nbTablesList[0].Name == "filter":
|
||||||
|
return IPTABLES
|
||||||
|
case err != nil:
|
||||||
|
log.Errorf("failed to list nftables tables on fw manager discovery: %s", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -176,15 +221,21 @@ func isIptablesClientAvailable(client *iptables.IPTables) bool {
|
|||||||
return err == nil
|
return err == nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// forceUserspaceFirewall reports whether the userspace firewall is forced.
|
||||||
|
// NB_FORCE_USERSPACE_ROUTER is an alias: forcing userspace routing implies the
|
||||||
|
// userspace firewall, since the two are no longer separable.
|
||||||
func forceUserspaceFirewall() bool {
|
func forceUserspaceFirewall() bool {
|
||||||
val := os.Getenv(EnvForceUserspaceFirewall)
|
return envForceBool(EnvForceUserspaceFirewall) || envForceBool(uspfilter.EnvForceUserspaceRouter)
|
||||||
|
}
|
||||||
|
|
||||||
|
func envForceBool(name string) bool {
|
||||||
|
val := os.Getenv(name)
|
||||||
if val == "" {
|
if val == "" {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
force, err := strconv.ParseBool(val)
|
force, err := strconv.ParseBool(val)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Warnf("failed to parse %s: %v", EnvForceUserspaceFirewall, err)
|
log.Warnf("failed to parse %s: %v", name, err)
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
return force
|
return force
|
||||||
|
|||||||
@@ -1,603 +0,0 @@
|
|||||||
package iptables
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"maps"
|
|
||||||
"net"
|
|
||||||
"slices"
|
|
||||||
|
|
||||||
"github.com/coreos/go-iptables/iptables"
|
|
||||||
"github.com/google/uuid"
|
|
||||||
ipset "github.com/lrh3321/ipset-go"
|
|
||||||
log "github.com/sirupsen/logrus"
|
|
||||||
|
|
||||||
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
|
||||||
"github.com/netbirdio/netbird/client/internal/statemanager"
|
|
||||||
nbnet "github.com/netbirdio/netbird/client/net"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
tableName = "filter"
|
|
||||||
|
|
||||||
// rules chains contains the effective ACL rules
|
|
||||||
chainNameInputRules = "NETBIRD-ACL-INPUT"
|
|
||||||
|
|
||||||
// mangleFwdKey is the entries map key for mangle FORWARD guard rules that prevent
|
|
||||||
// external DNAT from bypassing ACL rules.
|
|
||||||
mangleFwdKey = "MANGLE-FORWARD"
|
|
||||||
)
|
|
||||||
|
|
||||||
type aclEntries map[string][][]string
|
|
||||||
|
|
||||||
type entry struct {
|
|
||||||
spec []string
|
|
||||||
position int
|
|
||||||
}
|
|
||||||
|
|
||||||
type aclManager struct {
|
|
||||||
iptablesClient *iptables.IPTables
|
|
||||||
wgIface iFaceMapper
|
|
||||||
entries aclEntries
|
|
||||||
optionalEntries map[string][]entry
|
|
||||||
ipsetStore *ipsetStore
|
|
||||||
v6 bool
|
|
||||||
ipsetSupported bool
|
|
||||||
|
|
||||||
stateManager *statemanager.Manager
|
|
||||||
}
|
|
||||||
|
|
||||||
func newAclManager(iptablesClient *iptables.IPTables, wgIface iFaceMapper) (*aclManager, error) {
|
|
||||||
return &aclManager{
|
|
||||||
iptablesClient: iptablesClient,
|
|
||||||
wgIface: wgIface,
|
|
||||||
entries: make(map[string][][]string),
|
|
||||||
optionalEntries: make(map[string][]entry),
|
|
||||||
ipsetStore: newIpsetStore(),
|
|
||||||
v6: iptablesClient.Proto() == iptables.ProtocolIPv6,
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) init(stateManager *statemanager.Manager) error {
|
|
||||||
m.stateManager = stateManager
|
|
||||||
|
|
||||||
m.ipsetSupported = m.probeIPSetSupport()
|
|
||||||
|
|
||||||
m.seedInitialEntries()
|
|
||||||
m.seedInitialOptionalEntries()
|
|
||||||
|
|
||||||
if err := m.cleanChains(); err != nil {
|
|
||||||
return fmt.Errorf("clean chains: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.createDefaultChains(); err != nil {
|
|
||||||
return fmt.Errorf("create default chains: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.updateState()
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) AddPeerFiltering(
|
|
||||||
id []byte,
|
|
||||||
ip net.IP,
|
|
||||||
protocol firewall.Protocol,
|
|
||||||
sPort *firewall.Port,
|
|
||||||
dPort *firewall.Port,
|
|
||||||
action firewall.Action,
|
|
||||||
ipsetName string,
|
|
||||||
) ([]firewall.Rule, error) {
|
|
||||||
chain := chainNameInputRules
|
|
||||||
|
|
||||||
ipsetName = transformIPsetName(ipsetName, sPort, dPort, action)
|
|
||||||
if m.v6 && ipsetName != "" {
|
|
||||||
ipsetName += "-v6"
|
|
||||||
}
|
|
||||||
// When the kernel lacks the required ipset hash module, fall back to
|
|
||||||
// per-IP iptables rules (pre-0.68 behavior) so ACLs keep working instead
|
|
||||||
// of silently leaving the chain empty.
|
|
||||||
if ipsetName != "" && !m.ipsetSupported {
|
|
||||||
ipsetName = ""
|
|
||||||
}
|
|
||||||
proto := protoForFamily(protocol, m.v6)
|
|
||||||
specs := filterRuleSpecs(ip, proto, sPort, dPort, action, ipsetName)
|
|
||||||
|
|
||||||
mangleSpecs := slices.Clone(specs)
|
|
||||||
mangleSpecs = append(mangleSpecs,
|
|
||||||
"-i", m.wgIface.Name(),
|
|
||||||
"-m", "addrtype", "--dst-type", "LOCAL",
|
|
||||||
"-j", "MARK", "--set-xmark", fmt.Sprintf("%#x", nbnet.PreroutingFwmarkRedirected),
|
|
||||||
)
|
|
||||||
|
|
||||||
specs = append(specs, "-j", actionToStr(action))
|
|
||||||
if ipsetName != "" {
|
|
||||||
if ipList, ipsetExists := m.ipsetStore.ipset(ipsetName); ipsetExists {
|
|
||||||
if err := m.addToIPSet(ipsetName, ip); err != nil {
|
|
||||||
return nil, fmt.Errorf("add IP to ipset: %w", err)
|
|
||||||
}
|
|
||||||
// if ruleset already exists it means we already have the firewall rule
|
|
||||||
// so we need to update IPs in the ruleset and return new fw.Rule object for ACL manager.
|
|
||||||
ipList.addIP(ip.String())
|
|
||||||
return []firewall.Rule{&Rule{
|
|
||||||
ruleID: uuid.New().String(),
|
|
||||||
ipsetName: ipsetName,
|
|
||||||
ip: ip.String(),
|
|
||||||
chain: chain,
|
|
||||||
specs: specs,
|
|
||||||
v6: m.v6,
|
|
||||||
}}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.flushIPSet(ipsetName); err != nil {
|
|
||||||
if errors.Is(err, ipset.ErrSetNotExist) {
|
|
||||||
log.Debugf("flush ipset %s before use: %v", ipsetName, err)
|
|
||||||
} else {
|
|
||||||
log.Errorf("flush ipset %s before use: %v", ipsetName, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := m.createIPSet(ipsetName); err != nil {
|
|
||||||
return nil, fmt.Errorf("create ipset: %w", err)
|
|
||||||
}
|
|
||||||
if err := m.addToIPSet(ipsetName, ip); err != nil {
|
|
||||||
return nil, fmt.Errorf("add IP to ipset: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ipList := newIpList(ip.String())
|
|
||||||
m.ipsetStore.addIpList(ipsetName, ipList)
|
|
||||||
}
|
|
||||||
|
|
||||||
ok, err := m.iptablesClient.Exists(tableFilter, chain, specs...)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to check rule: %w", err)
|
|
||||||
}
|
|
||||||
if ok {
|
|
||||||
return nil, fmt.Errorf("rule already exists")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Insert DROP rules at the beginning, append ACCEPT rules at the end
|
|
||||||
if action == firewall.ActionDrop {
|
|
||||||
// Insert at the beginning of the chain (position 1)
|
|
||||||
err = m.iptablesClient.Insert(tableFilter, chain, 1, specs...)
|
|
||||||
} else {
|
|
||||||
err = m.iptablesClient.Append(tableFilter, chain, specs...)
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.iptablesClient.Append(tableMangle, chainRTPRE, mangleSpecs...); err != nil {
|
|
||||||
log.Errorf("failed to add mangle rule: %v", err)
|
|
||||||
mangleSpecs = nil
|
|
||||||
}
|
|
||||||
|
|
||||||
rule := &Rule{
|
|
||||||
ruleID: uuid.New().String(),
|
|
||||||
specs: specs,
|
|
||||||
mangleSpecs: mangleSpecs,
|
|
||||||
ipsetName: ipsetName,
|
|
||||||
ip: ip.String(),
|
|
||||||
chain: chain,
|
|
||||||
v6: m.v6,
|
|
||||||
}
|
|
||||||
|
|
||||||
m.updateState()
|
|
||||||
|
|
||||||
return []firewall.Rule{rule}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// DeletePeerRule from the firewall by rule definition
|
|
||||||
func (m *aclManager) DeletePeerRule(rule firewall.Rule) error {
|
|
||||||
r, ok := rule.(*Rule)
|
|
||||||
if !ok {
|
|
||||||
return fmt.Errorf("invalid rule type")
|
|
||||||
}
|
|
||||||
|
|
||||||
shouldDestroyIpset := false
|
|
||||||
if ipsetList, ok := m.ipsetStore.ipset(r.ipsetName); ok {
|
|
||||||
// delete IP from ruleset IPs list and ipset
|
|
||||||
if _, ok := ipsetList.ips[r.ip]; ok {
|
|
||||||
ip := net.ParseIP(r.ip)
|
|
||||||
if ip == nil {
|
|
||||||
return fmt.Errorf("parse IP %s", r.ip)
|
|
||||||
}
|
|
||||||
if err := m.delFromIPSet(r.ipsetName, ip); err != nil {
|
|
||||||
return fmt.Errorf("delete ip from ipset: %w", err)
|
|
||||||
}
|
|
||||||
delete(ipsetList.ips, r.ip)
|
|
||||||
}
|
|
||||||
|
|
||||||
// if after delete, set still contains other IPs,
|
|
||||||
// no need to delete firewall rule and we should exit here
|
|
||||||
if len(ipsetList.ips) != 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// we delete last IP from the set, that means we need to delete
|
|
||||||
// set itself and associated firewall rule too
|
|
||||||
m.ipsetStore.deleteIpset(r.ipsetName)
|
|
||||||
shouldDestroyIpset = true
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.iptablesClient.Delete(tableName, r.chain, r.specs...); err != nil {
|
|
||||||
return fmt.Errorf("failed to delete rule: %s, %v: %w", r.chain, r.specs, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if r.mangleSpecs != nil {
|
|
||||||
if err := m.iptablesClient.Delete(tableMangle, chainRTPRE, r.mangleSpecs...); err != nil {
|
|
||||||
log.Errorf("failed to delete mangle rule: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if shouldDestroyIpset {
|
|
||||||
if err := m.destroyIPSet(r.ipsetName); err != nil {
|
|
||||||
if errors.Is(err, ipset.ErrBusy) || errors.Is(err, ipset.ErrSetNotExist) {
|
|
||||||
log.Debugf("destroy empty ipset: %v", err)
|
|
||||||
} else {
|
|
||||||
log.Errorf("destroy empty ipset: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
m.updateState()
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) Reset() error {
|
|
||||||
if err := m.cleanChains(); err != nil {
|
|
||||||
return fmt.Errorf("clean chains: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.updateState()
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// todo write less destructive cleanup mechanism
|
|
||||||
func (m *aclManager) cleanChains() error {
|
|
||||||
ok, err := m.iptablesClient.ChainExists(tableName, chainNameInputRules)
|
|
||||||
if err != nil {
|
|
||||||
log.Debugf("failed to list chains: %s", err)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if ok {
|
|
||||||
for _, rule := range m.entries["INPUT"] {
|
|
||||||
err := m.iptablesClient.DeleteIfExists(tableName, "INPUT", rule...)
|
|
||||||
if err != nil {
|
|
||||||
log.Errorf("failed to delete rule: %v, %s", rule, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, rule := range m.entries["FORWARD"] {
|
|
||||||
err := m.iptablesClient.DeleteIfExists(tableName, "FORWARD", rule...)
|
|
||||||
if err != nil {
|
|
||||||
log.Errorf("failed to delete rule: %v, %s", rule, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
err = m.iptablesClient.ClearAndDeleteChain(tableName, chainNameInputRules)
|
|
||||||
if err != nil {
|
|
||||||
log.Debugf("failed to clear and delete %s chain: %s", chainNameInputRules, err)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
ok, err = m.iptablesClient.ChainExists("mangle", "PREROUTING")
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("list chains: %w", err)
|
|
||||||
}
|
|
||||||
if ok {
|
|
||||||
for _, rule := range m.entries["PREROUTING"] {
|
|
||||||
err := m.iptablesClient.DeleteIfExists("mangle", "PREROUTING", rule...)
|
|
||||||
if err != nil {
|
|
||||||
log.Errorf("failed to delete rule: %v, %s", rule, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, rule := range m.entries[mangleFwdKey] {
|
|
||||||
if err := m.iptablesClient.DeleteIfExists(tableMangle, chainFORWARD, rule...); err != nil {
|
|
||||||
log.Errorf("failed to delete mangle FORWARD guard rule: %v, %s", rule, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, ipsetName := range m.ipsetStore.ipsetNames() {
|
|
||||||
if err := m.flushIPSet(ipsetName); err != nil {
|
|
||||||
if errors.Is(err, ipset.ErrSetNotExist) {
|
|
||||||
log.Debugf("flush ipset %q during reset: %v", ipsetName, err)
|
|
||||||
} else {
|
|
||||||
log.Errorf("flush ipset %q during reset: %v", ipsetName, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := m.destroyIPSet(ipsetName); err != nil {
|
|
||||||
if errors.Is(err, ipset.ErrBusy) || errors.Is(err, ipset.ErrSetNotExist) {
|
|
||||||
log.Debugf("destroy ipset %q during reset: %v", ipsetName, err)
|
|
||||||
} else {
|
|
||||||
log.Errorf("destroy ipset %q during reset: %v", ipsetName, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
m.ipsetStore.deleteIpset(ipsetName)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) createDefaultChains() error {
|
|
||||||
// chain netbird-acl-input-rules
|
|
||||||
if err := m.iptablesClient.NewChain(tableName, chainNameInputRules); err != nil {
|
|
||||||
log.Debugf("failed to create '%s' chain: %s", chainNameInputRules, err)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
for chainName, rules := range m.entries {
|
|
||||||
// mangle FORWARD guard rules are handled separately below
|
|
||||||
if chainName == mangleFwdKey {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
for _, rule := range rules {
|
|
||||||
if err := m.iptablesClient.InsertUnique(tableName, chainName, 1, rule...); err != nil {
|
|
||||||
log.Debugf("failed to create input chain jump rule: %s", err)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for chainName, entries := range m.optionalEntries {
|
|
||||||
for _, entry := range entries {
|
|
||||||
if err := m.iptablesClient.InsertUnique(tableName, chainName, entry.position, entry.spec...); err != nil {
|
|
||||||
log.Errorf("failed to insert optional entry %v: %v", entry.spec, err)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
m.entries[chainName] = append(m.entries[chainName], entry.spec)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
clear(m.optionalEntries)
|
|
||||||
|
|
||||||
// Insert mangle FORWARD guard rules to prevent external DNAT bypass.
|
|
||||||
for _, rule := range m.entries[mangleFwdKey] {
|
|
||||||
if err := m.iptablesClient.AppendUnique(tableMangle, chainFORWARD, rule...); err != nil {
|
|
||||||
log.Errorf("failed to add mangle FORWARD guard rule: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// seedInitialEntries adds default rules to the entries map, rules are inserted on pos 1, hence the order is reversed.
|
|
||||||
// We want to make sure our traffic is not dropped by existing rules.
|
|
||||||
|
|
||||||
// The existing FORWARD rules/policies decide outbound traffic towards our interface.
|
|
||||||
// In case the FORWARD policy is set to "drop", we add an established/related rule to allow return traffic for the inbound rule.
|
|
||||||
func (m *aclManager) seedInitialEntries() {
|
|
||||||
established := getConntrackEstablished()
|
|
||||||
|
|
||||||
m.appendToEntries("INPUT", []string{"-i", m.wgIface.Name(), "-j", "DROP"})
|
|
||||||
m.appendToEntries("INPUT", []string{"-i", m.wgIface.Name(), "-j", chainNameInputRules})
|
|
||||||
m.appendToEntries("INPUT", append([]string{"-i", m.wgIface.Name()}, established...))
|
|
||||||
|
|
||||||
// Inbound is handled by our ACLs, the rest is dropped.
|
|
||||||
// For outbound we respect the FORWARD policy. However, we need to allow established/related traffic for inbound rules.
|
|
||||||
m.appendToEntries("FORWARD", []string{"-i", m.wgIface.Name(), "-j", "DROP"})
|
|
||||||
|
|
||||||
m.appendToEntries("FORWARD", []string{"-o", m.wgIface.Name(), "-j", chainRTFWDOUT})
|
|
||||||
m.appendToEntries("FORWARD", []string{"-i", m.wgIface.Name(), "-j", chainRTFWDIN})
|
|
||||||
|
|
||||||
// Mangle FORWARD guard: when external DNAT redirects traffic from the wg interface, it
|
|
||||||
// traverses FORWARD instead of INPUT, bypassing ACL rules. ACCEPT rules in filter FORWARD
|
|
||||||
// can be inserted above ours. Mangle runs before filter, so these guard rules enforce the
|
|
||||||
// ACL mark check where it cannot be overridden.
|
|
||||||
m.appendToEntries(mangleFwdKey, []string{
|
|
||||||
"-i", m.wgIface.Name(),
|
|
||||||
"-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED",
|
|
||||||
"-j", "ACCEPT",
|
|
||||||
})
|
|
||||||
m.appendToEntries(mangleFwdKey, []string{
|
|
||||||
"-i", m.wgIface.Name(),
|
|
||||||
"-m", "conntrack", "--ctstate", "DNAT",
|
|
||||||
"-m", "mark", "!", "--mark", fmt.Sprintf("%#x", nbnet.PreroutingFwmarkRedirected),
|
|
||||||
"-j", "DROP",
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) seedInitialOptionalEntries() {
|
|
||||||
m.optionalEntries["FORWARD"] = []entry{
|
|
||||||
{
|
|
||||||
spec: []string{"-m", "mark", "--mark", fmt.Sprintf("%#x", nbnet.PreroutingFwmarkRedirected), "-j", "ACCEPT"},
|
|
||||||
position: 2,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) appendToEntries(chainName string, spec []string) {
|
|
||||||
m.entries[chainName] = append(m.entries[chainName], spec)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) updateState() {
|
|
||||||
if m.stateManager == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
var currentState *ShutdownState
|
|
||||||
if existing := m.stateManager.GetState(currentState); existing != nil {
|
|
||||||
if existingState, ok := existing.(*ShutdownState); ok {
|
|
||||||
currentState = existingState
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if currentState == nil {
|
|
||||||
currentState = &ShutdownState{}
|
|
||||||
}
|
|
||||||
|
|
||||||
currentState.Lock()
|
|
||||||
defer currentState.Unlock()
|
|
||||||
|
|
||||||
// Clone the maps so the persisted state holds a private snapshot. The
|
|
||||||
// live maps keep being mutated by subsequent rule operations while the
|
|
||||||
// state manager marshals the state from its periodic-save goroutine.
|
|
||||||
// Sharing them by reference races the two and aborts the process with a
|
|
||||||
// concurrent map iteration and write.
|
|
||||||
if m.v6 {
|
|
||||||
currentState.ACLEntries6 = maps.Clone(m.entries)
|
|
||||||
currentState.ACLIPsetStore6 = m.ipsetStore.clone()
|
|
||||||
} else {
|
|
||||||
currentState.ACLEntries = maps.Clone(m.entries)
|
|
||||||
currentState.ACLIPsetStore = m.ipsetStore.clone()
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.stateManager.UpdateState(currentState); err != nil {
|
|
||||||
log.Errorf("failed to update state: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// filterRuleSpecs returns the specs of a filtering rule
|
|
||||||
// protoForFamily translates ICMP to ICMPv6 for ip6tables.
|
|
||||||
// ip6tables requires "ipv6-icmp" (or "icmpv6") instead of "icmp".
|
|
||||||
func protoForFamily(protocol firewall.Protocol, v6 bool) string {
|
|
||||||
if v6 && protocol == firewall.ProtocolICMP {
|
|
||||||
return "ipv6-icmp"
|
|
||||||
}
|
|
||||||
return string(protocol)
|
|
||||||
}
|
|
||||||
|
|
||||||
func filterRuleSpecs(ip net.IP, protocol string, sPort, dPort *firewall.Port, action firewall.Action, ipsetName string) (specs []string) {
|
|
||||||
// don't use IP matching if IP is 0.0.0.0
|
|
||||||
matchByIP := !ip.IsUnspecified()
|
|
||||||
|
|
||||||
if matchByIP {
|
|
||||||
if ipsetName != "" {
|
|
||||||
specs = append(specs, "-m", "set", "--match-set", ipsetName, "src")
|
|
||||||
} else {
|
|
||||||
specs = append(specs, "-s", ip.String())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if protocol != "all" {
|
|
||||||
specs = append(specs, "-p", protocol)
|
|
||||||
}
|
|
||||||
specs = append(specs, applyPort("--sport", sPort)...)
|
|
||||||
specs = append(specs, applyPort("--dport", dPort)...)
|
|
||||||
return specs
|
|
||||||
}
|
|
||||||
|
|
||||||
func actionToStr(action firewall.Action) string {
|
|
||||||
if action == firewall.ActionAccept {
|
|
||||||
return "ACCEPT"
|
|
||||||
}
|
|
||||||
return "DROP"
|
|
||||||
}
|
|
||||||
|
|
||||||
func transformIPsetName(ipsetName string, sPort, dPort *firewall.Port, action firewall.Action) string {
|
|
||||||
if ipsetName == "" {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
actionSuffix := ""
|
|
||||||
if action == firewall.ActionDrop {
|
|
||||||
actionSuffix = "-drop"
|
|
||||||
}
|
|
||||||
|
|
||||||
switch {
|
|
||||||
case sPort != nil && dPort != nil:
|
|
||||||
return ipsetName + "-sport-dport" + actionSuffix
|
|
||||||
case sPort != nil:
|
|
||||||
return ipsetName + "-sport" + actionSuffix
|
|
||||||
case dPort != nil:
|
|
||||||
return ipsetName + "-dport" + actionSuffix
|
|
||||||
default:
|
|
||||||
return ipsetName + actionSuffix
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// probeIPSetSupport checks whether the kernel can create the ipset type used for
|
|
||||||
// ACL rules. On kernels lacking the required ipset hash module, ipset creation
|
|
||||||
// fails (e.g. "invalid argument"), which would otherwise leave the ACL chain
|
|
||||||
// empty and silently drop all policy-permitted inbound traffic. When unsupported,
|
|
||||||
// the manager falls back to per-IP iptables rules.
|
|
||||||
func (m *aclManager) probeIPSetSupport() bool {
|
|
||||||
// Use a unique name so concurrent processes don't collide and we only ever
|
|
||||||
// destroy the set we created ourselves. ipset names are limited to 31 chars,
|
|
||||||
// so use a short random suffix.
|
|
||||||
probeName := "nb-probe-" + uuid.New().String()[:8]
|
|
||||||
|
|
||||||
opts := ipset.CreateOptions{
|
|
||||||
Replace: true,
|
|
||||||
}
|
|
||||||
if m.v6 {
|
|
||||||
opts.Family = ipset.FamilyIPV6
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := ipset.Create(probeName, ipset.TypeHashNet, opts); err != nil {
|
|
||||||
log.Warnf("ipset is not available (failed to create probe set: %v); "+
|
|
||||||
"falling back to per-IP iptables ACL rules. Ensure the kernel provides "+
|
|
||||||
"the ipset hash:net module (ip_set_hash_net) for better performance with large rule sets", err)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
defer func() {
|
|
||||||
if err := ipset.Destroy(probeName); err != nil {
|
|
||||||
log.Debugf("destroy ipset probe set %q: %v", probeName, err)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) createIPSet(name string) error {
|
|
||||||
opts := ipset.CreateOptions{
|
|
||||||
Replace: true,
|
|
||||||
}
|
|
||||||
if m.v6 {
|
|
||||||
opts.Family = ipset.FamilyIPV6
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := ipset.Create(name, ipset.TypeHashNet, opts); err != nil {
|
|
||||||
return fmt.Errorf("create ipset %s: %w", name, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Debugf("created ipset %s with type hash:net", name)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) addToIPSet(name string, ip net.IP) error {
|
|
||||||
cidr := uint8(32)
|
|
||||||
if ip.To4() == nil {
|
|
||||||
cidr = 128
|
|
||||||
}
|
|
||||||
|
|
||||||
entry := &ipset.Entry{
|
|
||||||
IP: ip,
|
|
||||||
CIDR: cidr,
|
|
||||||
Replace: true,
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := ipset.Add(name, entry); err != nil {
|
|
||||||
return fmt.Errorf("add IP to ipset %s: %w", name, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) delFromIPSet(name string, ip net.IP) error {
|
|
||||||
cidr := uint8(32)
|
|
||||||
if ip.To4() == nil {
|
|
||||||
cidr = 128
|
|
||||||
}
|
|
||||||
|
|
||||||
entry := &ipset.Entry{
|
|
||||||
IP: ip,
|
|
||||||
CIDR: cidr,
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := ipset.Del(name, entry); err != nil {
|
|
||||||
return fmt.Errorf("delete IP from ipset %s: %w", name, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) flushIPSet(name string) error {
|
|
||||||
return ipset.Flush(name)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *aclManager) destroyIPSet(name string) error {
|
|
||||||
return ipset.Destroy(name)
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,346 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package iptables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
nbnet "github.com/netbirdio/netbird/client/net"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (r *family) createContainers() error {
|
||||||
|
for _, chainInfo := range []struct {
|
||||||
|
chain string
|
||||||
|
table string
|
||||||
|
}{
|
||||||
|
{chainRTFwdIn, tableFilter},
|
||||||
|
{chainRTFwdOut, tableFilter},
|
||||||
|
{chainRTPre, tableMangle},
|
||||||
|
{chainRTNAT, tableNat},
|
||||||
|
{chainRTRdr, tableNat},
|
||||||
|
{chainRTMSSClamp, tableMangle},
|
||||||
|
} {
|
||||||
|
// Fallback: clear chains that survived an unclean shutdown.
|
||||||
|
if ok, _ := r.iptablesClient.ChainExists(chainInfo.table, chainInfo.chain); ok {
|
||||||
|
if err := r.iptablesClient.ClearAndDeleteChain(chainInfo.table, chainInfo.chain); err != nil {
|
||||||
|
log.Warnf("clear stale chain %s in %s: %v", chainInfo.chain, chainInfo.table, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := r.iptablesClient.NewChain(chainInfo.table, chainInfo.chain); err != nil {
|
||||||
|
return fmt.Errorf("create chain %s in table %s: %w", chainInfo.chain, chainInfo.table, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.insertEstablishedRule(chainRTFwdIn); err != nil {
|
||||||
|
return fmt.Errorf("insert established rule: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.insertEstablishedRule(chainRTFwdOut); err != nil {
|
||||||
|
return fmt.Errorf("insert established rule: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.addPostroutingRules(); err != nil {
|
||||||
|
return fmt.Errorf("add static nat rules: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.addJumpRules(); err != nil {
|
||||||
|
return fmt.Errorf("add jump rules: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.addMSSClampingRules(); err != nil {
|
||||||
|
log.Errorf("failed to add MSS clamping rules: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addJumpRules() error {
|
||||||
|
// Jump to nat chain
|
||||||
|
natRule := jumpRuleSpec(chainRTNAT)
|
||||||
|
if err := r.iptablesClient.Insert(tableNat, chainPostrouting, 1, natRule...); err != nil {
|
||||||
|
return fmt.Errorf("add nat postrouting jump rule: %w", err)
|
||||||
|
}
|
||||||
|
r.rules[jumpNATPost] = natRule
|
||||||
|
|
||||||
|
// Jump to mangle prerouting chain
|
||||||
|
preRule := jumpRuleSpec(chainRTPre)
|
||||||
|
if err := r.iptablesClient.Insert(tableMangle, chainPrerouting, 1, preRule...); err != nil {
|
||||||
|
return fmt.Errorf("add mangle prerouting jump rule: %w", err)
|
||||||
|
}
|
||||||
|
r.rules[jumpManglePre] = preRule
|
||||||
|
|
||||||
|
// Jump to nat prerouting chain
|
||||||
|
rdrRule := jumpRuleSpec(chainRTRdr)
|
||||||
|
if err := r.iptablesClient.Insert(tableNat, chainPrerouting, 1, rdrRule...); err != nil {
|
||||||
|
return fmt.Errorf("add nat prerouting jump rule: %w", err)
|
||||||
|
}
|
||||||
|
r.rules[jumpNATPre] = rdrRule
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) setupDataPlaneMark() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
preRule := []string{
|
||||||
|
"-i", r.wgIface.Name(),
|
||||||
|
"-m", "conntrack", "--ctstate", "NEW",
|
||||||
|
"-j", "CONNMARK", "--set-mark", fmt.Sprintf("%#x", nbnet.DataPlaneMarkIn),
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.iptablesClient.AppendUnique(tableMangle, chainPrerouting, preRule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("add mangle prerouting rule: %w", err))
|
||||||
|
} else {
|
||||||
|
r.rules[markManglePre] = preRule
|
||||||
|
}
|
||||||
|
|
||||||
|
postRule := []string{
|
||||||
|
"-o", r.wgIface.Name(),
|
||||||
|
"-m", "conntrack", "--ctstate", "NEW",
|
||||||
|
"-j", "CONNMARK", "--set-mark", fmt.Sprintf("%#x", nbnet.DataPlaneMarkOut),
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.iptablesClient.AppendUnique(tableMangle, chainPostrouting, postRule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("add mangle postrouting rule: %w", err))
|
||||||
|
} else {
|
||||||
|
r.rules[markManglePost] = postRule
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// seedInitialEntries adds default rules to the entries map. Rules are
|
||||||
|
// inserted at position 1, so the order here is reversed.
|
||||||
|
//
|
||||||
|
// Existing FORWARD policy decides outbound traffic towards our
|
||||||
|
// interface. If FORWARD policy is "drop", we add an
|
||||||
|
// established/related rule to allow return traffic for inbound rules.
|
||||||
|
func (r *family) seedInitialEntries() {
|
||||||
|
established := getConntrackEstablished()
|
||||||
|
|
||||||
|
r.appendToEntries(chainInput, []string{"-i", r.wgIface.Name(), "-j", "DROP"})
|
||||||
|
r.appendToEntries(chainInput, []string{"-i", r.wgIface.Name(), "-j", chainACLInput})
|
||||||
|
r.appendToEntries(chainInput, append([]string{"-i", r.wgIface.Name()}, established...))
|
||||||
|
|
||||||
|
r.appendToEntries(chainForward, []string{"-i", r.wgIface.Name(), "-j", "DROP"})
|
||||||
|
r.appendToEntries(chainForward, []string{"-o", r.wgIface.Name(), "-j", chainRTFwdOut})
|
||||||
|
r.appendToEntries(chainForward, []string{"-i", r.wgIface.Name(), "-j", chainRTFwdIn})
|
||||||
|
|
||||||
|
// Mangle FORWARD guard: when external DNAT redirects traffic from
|
||||||
|
// the wg interface, it traverses FORWARD instead of INPUT,
|
||||||
|
// bypassing ACL rules. ACCEPT rules in filter FORWARD can be
|
||||||
|
// inserted above ours. Mangle runs before filter, so these guard
|
||||||
|
// rules enforce the ACL mark check where it cannot be overridden.
|
||||||
|
r.appendToEntries(mangleForwardKey, []string{
|
||||||
|
"-i", r.wgIface.Name(),
|
||||||
|
"-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED",
|
||||||
|
"-j", "ACCEPT",
|
||||||
|
})
|
||||||
|
r.appendToEntries(mangleForwardKey, []string{
|
||||||
|
"-i", r.wgIface.Name(),
|
||||||
|
"-m", "conntrack", "--ctstate", "DNAT",
|
||||||
|
"-m", "mark", "!", "--mark", fmt.Sprintf("%#x", nbnet.PreroutingFwmarkRedirected),
|
||||||
|
"-j", "DROP",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) seedInitialOptionalEntries() {
|
||||||
|
r.optionalEntries[chainForward] = []entry{
|
||||||
|
{
|
||||||
|
spec: []string{"-m", "mark", "--mark", fmt.Sprintf("%#x", nbnet.PreroutingFwmarkRedirected), "-j", "ACCEPT"},
|
||||||
|
position: 2,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) appendToEntries(chain chainKey, spec ruleSpec) {
|
||||||
|
r.entries[chain] = append(r.entries[chain], spec)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) createDefaultChains() error {
|
||||||
|
if err := r.iptablesClient.NewChain(tableFilter, chainACLInput); err != nil {
|
||||||
|
return fmt.Errorf("create %s chain: %w", chainACLInput, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for chain, rules := range r.entries {
|
||||||
|
// mangle FORWARD guard rules are handled separately below
|
||||||
|
if chain == mangleForwardKey {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, rule := range rules {
|
||||||
|
if err := r.iptablesClient.InsertUnique(tableFilter, string(chain), 1, rule...); err != nil {
|
||||||
|
return fmt.Errorf("insert jump rule into %s: %w", chain, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for chain, entries := range r.optionalEntries {
|
||||||
|
for _, entry := range entries {
|
||||||
|
if err := r.iptablesClient.InsertUnique(tableFilter, string(chain), entry.position, entry.spec...); err != nil {
|
||||||
|
log.Errorf("failed to insert optional entry %v: %v", entry.spec, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
r.entries[chain] = append(r.entries[chain], entry.spec)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
clear(r.optionalEntries)
|
||||||
|
|
||||||
|
// Insert mangle FORWARD guard rules to prevent external DNAT bypass.
|
||||||
|
for _, rule := range r.entries[mangleForwardKey] {
|
||||||
|
if err := r.iptablesClient.AppendUnique(tableMangle, chainForward, rule...); err != nil {
|
||||||
|
log.Errorf("failed to add mangle FORWARD guard rule: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) cleanUpDefaultForwardRules() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
|
||||||
|
// cleanJumpRules removes the OUTPUT jump to NETBIRD-NAT-OUTPUT among
|
||||||
|
// the others, so the chain below deletes cleanly instead of failing
|
||||||
|
// with "device or resource busy".
|
||||||
|
if err := r.cleanJumpRules(); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("clean jump rules: %w", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, chainInfo := range []struct {
|
||||||
|
chain string
|
||||||
|
table string
|
||||||
|
}{
|
||||||
|
{chainRTFwdIn, tableFilter},
|
||||||
|
{chainRTFwdOut, tableFilter},
|
||||||
|
{chainRTPre, tableMangle},
|
||||||
|
{chainRTNAT, tableNat},
|
||||||
|
{chainRTRdr, tableNat},
|
||||||
|
{chainNATOutput, tableNat},
|
||||||
|
{chainRTMSSClamp, tableMangle},
|
||||||
|
} {
|
||||||
|
ok, err := r.iptablesClient.ChainExists(chainInfo.table, chainInfo.chain)
|
||||||
|
if err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("check chain %s in table %s: %w", chainInfo.chain, chainInfo.table, err))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if ok {
|
||||||
|
if err := r.iptablesClient.ClearAndDeleteChain(chainInfo.table, chainInfo.chain); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("clear and delete chain %s in table %s: %w", chainInfo.chain, chainInfo.table, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) cleanJumpRules() error {
|
||||||
|
// locations maps each jump rule to the built-in table and chain it
|
||||||
|
// was inserted into, plus the netbird chain it targets.
|
||||||
|
locations := map[firewall.RuleID]struct{ table, chain, target string }{
|
||||||
|
jumpNATPost: {tableNat, chainPostrouting, chainRTNAT},
|
||||||
|
jumpManglePre: {tableMangle, chainPrerouting, chainRTPre},
|
||||||
|
jumpNATPre: {tableNat, chainPrerouting, chainRTRdr},
|
||||||
|
jumpMSSClamp: {tableMangle, chainForward, chainRTMSSClamp},
|
||||||
|
jumpNATOutput: {tableNat, chainOutput, chainNATOutput},
|
||||||
|
}
|
||||||
|
|
||||||
|
var merr *multierror.Error
|
||||||
|
for ruleID, loc := range locations {
|
||||||
|
rule, exists := r.rules[ruleID]
|
||||||
|
if !exists {
|
||||||
|
// Untracked (e.g. fresh start after an unclean shutdown with no
|
||||||
|
// restored state): if the target chain survived, remove the stale
|
||||||
|
// jump to it so the chain can be deleted.
|
||||||
|
ok, err := r.iptablesClient.ChainExists(loc.table, loc.target)
|
||||||
|
if err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("check chain %s in table %s: %w", loc.target, loc.table, err))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
rule = jumpRuleSpec(loc.target)
|
||||||
|
}
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(loc.table, loc.chain, rule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete rule from chain %s in table %s: %w", loc.chain, loc.table, err))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
}
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// jumpRuleSpec builds the iptables rule spec that jumps to target. Create
|
||||||
|
// and cleanup sites share it so the installed and deleted specs cannot drift.
|
||||||
|
func jumpRuleSpec(target string) []string {
|
||||||
|
return []string{"-j", target}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) cleanAclChains() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
|
||||||
|
if err := r.cleanInputAclChain(); err != nil {
|
||||||
|
merr = multierror.Append(merr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, rule := range r.entries[mangleForwardKey] {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableMangle, chainForward, rule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete mangle %s guard rule %v: %w", chainForward, rule, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) cleanInputAclChain() error {
|
||||||
|
ok, err := r.iptablesClient.ChainExists(tableFilter, chainACLInput)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("check chain %s: %w", chainACLInput, err)
|
||||||
|
}
|
||||||
|
if !ok {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var merr *multierror.Error
|
||||||
|
for _, rule := range r.entries[chainInput] {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableFilter, chainInput, rule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete %s rule %v: %w", chainInput, rule, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, rule := range r.entries[chainForward] {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableFilter, chainForward, rule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete %s rule %v: %w", chainForward, rule, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.iptablesClient.ClearAndDeleteChain(tableFilter, chainACLInput); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("clear and delete %s chain: %w", chainACLInput, err))
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) cleanupDataPlaneMark() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
if preRule, exists := r.rules[markManglePre]; exists {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableMangle, chainPrerouting, preRule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove mangle prerouting rule: %w", err))
|
||||||
|
} else {
|
||||||
|
delete(r.rules, markManglePre)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if postRule, exists := r.rules[markManglePost]; exists {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableMangle, chainPostrouting, postRule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove mangle postrouting rule: %w", err))
|
||||||
|
} else {
|
||||||
|
delete(r.rules, markManglePost)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
@@ -0,0 +1,302 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package iptables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (r *family) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) {
|
||||||
|
ruleID := rule.ID()
|
||||||
|
if _, exists := r.rules[ruleID+dnatSuffix]; exists {
|
||||||
|
return rule, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
toDestination := rule.TranslatedAddress.String()
|
||||||
|
switch {
|
||||||
|
case len(rule.TranslatedPort.Values) == 0:
|
||||||
|
// no translated port, use original port
|
||||||
|
case len(rule.TranslatedPort.Values) == 1:
|
||||||
|
toDestination += fmt.Sprintf(":%d", rule.TranslatedPort.Values[0])
|
||||||
|
case rule.TranslatedPort.IsRange && len(rule.TranslatedPort.Values) == 2:
|
||||||
|
// need the "/originalport" suffix to avoid dnat port randomization
|
||||||
|
toDestination += fmt.Sprintf(":%d-%d/%d", rule.TranslatedPort.Values[0], rule.TranslatedPort.Values[1], rule.DestinationPort.Values[0])
|
||||||
|
default:
|
||||||
|
return nil, fmt.Errorf("invalid translated port: %v", rule.TranslatedPort)
|
||||||
|
}
|
||||||
|
|
||||||
|
proto := strings.ToLower(string(rule.Protocol))
|
||||||
|
|
||||||
|
rules := make(map[firewall.RuleID]ruleInfo, 3)
|
||||||
|
|
||||||
|
// DNAT rule
|
||||||
|
dnatRule := []string{
|
||||||
|
"!", "-i", r.wgIface.Name(),
|
||||||
|
"-p", proto,
|
||||||
|
"-j", "DNAT",
|
||||||
|
"--to-destination", toDestination,
|
||||||
|
}
|
||||||
|
dnatRule = append(dnatRule, applyPort("--dport", &rule.DestinationPort)...)
|
||||||
|
rules[ruleID+dnatSuffix] = ruleInfo{
|
||||||
|
table: tableNat,
|
||||||
|
chain: chainRTRdr,
|
||||||
|
rule: dnatRule,
|
||||||
|
}
|
||||||
|
|
||||||
|
// SNAT rule
|
||||||
|
snatRule := []string{
|
||||||
|
"-o", r.wgIface.Name(),
|
||||||
|
"-p", proto,
|
||||||
|
"-d", rule.TranslatedAddress.String(),
|
||||||
|
"-j", "MASQUERADE",
|
||||||
|
}
|
||||||
|
snatRule = append(snatRule, applyPort("--dport", &rule.TranslatedPort)...)
|
||||||
|
rules[ruleID+snatSuffix] = ruleInfo{
|
||||||
|
table: tableNat,
|
||||||
|
chain: chainRTNAT,
|
||||||
|
rule: snatRule,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Forward filtering rule, if fwd policy is DROP
|
||||||
|
forwardRule := []string{
|
||||||
|
"-o", r.wgIface.Name(),
|
||||||
|
"-p", proto,
|
||||||
|
"-d", rule.TranslatedAddress.String(),
|
||||||
|
"-j", "ACCEPT",
|
||||||
|
}
|
||||||
|
forwardRule = append(forwardRule, applyPort("--dport", &rule.TranslatedPort)...)
|
||||||
|
rules[ruleID+fwdSuffix] = ruleInfo{
|
||||||
|
table: tableFilter,
|
||||||
|
chain: chainRTFwdOut,
|
||||||
|
rule: forwardRule,
|
||||||
|
}
|
||||||
|
|
||||||
|
for key, ruleInfo := range rules {
|
||||||
|
if err := r.iptablesClient.Append(ruleInfo.table, ruleInfo.chain, ruleInfo.rule...); err != nil {
|
||||||
|
r.cleanupFailedDNATAdd(rules)
|
||||||
|
return nil, fmt.Errorf("add rule %s: %w", key, err)
|
||||||
|
}
|
||||||
|
r.rules[key] = ruleInfo.rule
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.ipFwdState.RequestForwarding(r.v6); err != nil {
|
||||||
|
r.cleanupFailedDNATAdd(rules)
|
||||||
|
return nil, fmt.Errorf("enable forwarding: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.updateState()
|
||||||
|
return rule, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// cleanupFailedDNATAdd removes the bookkeeping written by a partially applied
|
||||||
|
// AddDNATRule before rolling back the kernel rules, so no entries remain that
|
||||||
|
// never got a forwarding refcount. rollbackRules re-adds entries it failed to
|
||||||
|
// remove from the kernel.
|
||||||
|
func (r *family) cleanupFailedDNATAdd(rules map[firewall.RuleID]ruleInfo) {
|
||||||
|
for key := range rules {
|
||||||
|
delete(r.rules, key)
|
||||||
|
}
|
||||||
|
if err := r.rollbackRules(rules); err != nil {
|
||||||
|
log.Errorf("rollback failed: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) rollbackRules(rules map[firewall.RuleID]ruleInfo) error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
for key, ruleInfo := range rules {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(ruleInfo.table, ruleInfo.chain, ruleInfo.rule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("rollback rule %s: %w", key, err))
|
||||||
|
// On rollback error, add to rules map for next cleanup
|
||||||
|
r.rules[key] = ruleInfo.rule
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if merr != nil {
|
||||||
|
r.updateState()
|
||||||
|
}
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) DeleteDNATRule(rule firewall.Rule) error {
|
||||||
|
ruleID := rule.ID()
|
||||||
|
|
||||||
|
_, hadDNAT := r.rules[ruleID+dnatSuffix]
|
||||||
|
_, hadSNAT := r.rules[ruleID+snatSuffix]
|
||||||
|
_, hadFWD := r.rules[ruleID+fwdSuffix]
|
||||||
|
if !hadDNAT && !hadSNAT && !hadFWD {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var merr *multierror.Error
|
||||||
|
if dnatRule, exists := r.rules[ruleID+dnatSuffix]; exists {
|
||||||
|
if err := r.iptablesClient.Delete(tableNat, chainRTRdr, dnatRule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete DNAT rule: %w", err))
|
||||||
|
} else {
|
||||||
|
delete(r.rules, ruleID+dnatSuffix)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if snatRule, exists := r.rules[ruleID+snatSuffix]; exists {
|
||||||
|
if err := r.iptablesClient.Delete(tableNat, chainRTNAT, snatRule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete SNAT rule: %w", err))
|
||||||
|
} else {
|
||||||
|
delete(r.rules, ruleID+snatSuffix)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if fwdRule, exists := r.rules[ruleID+fwdSuffix]; exists {
|
||||||
|
if err := r.iptablesClient.Delete(tableFilter, chainRTFwdOut, fwdRule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete forward rule: %w", err))
|
||||||
|
} else {
|
||||||
|
delete(r.rules, ruleID+fwdSuffix)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Release the refcount only once all rules are gone from the kernel. On
|
||||||
|
// partial failure the failed entries stay in r.rules so a retry can remove
|
||||||
|
// them and release then.
|
||||||
|
if merr == nil {
|
||||||
|
r.releaseForwarding()
|
||||||
|
}
|
||||||
|
|
||||||
|
r.updateState()
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// releaseForwarding drops one IP forwarding reference, logging any error.
|
||||||
|
func (r *family) releaseForwarding() {
|
||||||
|
if err := r.ipFwdState.ReleaseForwarding(r.v6); err != nil {
|
||||||
|
log.Errorf("release IP forwarding: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) AddInboundDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
||||||
|
ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
||||||
|
|
||||||
|
if _, exists := r.rules[ruleID]; exists {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
dnatRule := []string{
|
||||||
|
"-i", r.wgIface.Name(),
|
||||||
|
"-p", strings.ToLower(protoForFamily(protocol, r.v6)),
|
||||||
|
"--dport", strconv.Itoa(int(originalPort)),
|
||||||
|
"-d", localAddr.String(),
|
||||||
|
"-m", "addrtype", "--dst-type", "LOCAL",
|
||||||
|
"-j", "DNAT",
|
||||||
|
"--to-destination", ":" + strconv.Itoa(int(translatedPort)),
|
||||||
|
}
|
||||||
|
|
||||||
|
info := ruleInfo{
|
||||||
|
table: tableNat,
|
||||||
|
chain: chainRTRdr,
|
||||||
|
rule: dnatRule,
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.iptablesClient.Append(info.table, info.chain, info.rule...); err != nil {
|
||||||
|
return fmt.Errorf("add inbound DNAT rule: %w", err)
|
||||||
|
}
|
||||||
|
r.rules[ruleID] = info.rule
|
||||||
|
|
||||||
|
r.updateState()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RemoveInboundDNAT removes an inbound DNAT rule.
|
||||||
|
func (r *family) RemoveInboundDNAT(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))
|
||||||
|
|
||||||
|
if dnatRule, exists := r.rules[ruleID]; exists {
|
||||||
|
if err := r.iptablesClient.Delete(tableNat, chainRTRdr, dnatRule...); err != nil {
|
||||||
|
return fmt.Errorf("delete inbound DNAT rule: %w", err)
|
||||||
|
}
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.updateState()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ensureNATOutputChain lazily creates the OUTPUT NAT chain and jump rule on first use.
|
||||||
|
func (r *family) ensureNATOutputChain() error {
|
||||||
|
if _, exists := r.rules[jumpNATOutput]; exists {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
chainExists, err := r.iptablesClient.ChainExists(tableNat, chainNATOutput)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("check chain %s: %w", chainNATOutput, err)
|
||||||
|
}
|
||||||
|
if !chainExists {
|
||||||
|
if err := r.iptablesClient.NewChain(tableNat, chainNATOutput); err != nil {
|
||||||
|
return fmt.Errorf("create chain %s: %w", chainNATOutput, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
jumpRule := jumpRuleSpec(chainNATOutput)
|
||||||
|
if err := r.iptablesClient.Insert(tableNat, chainOutput, 1, jumpRule...); err != nil {
|
||||||
|
if !chainExists {
|
||||||
|
if delErr := r.iptablesClient.ClearAndDeleteChain(tableNat, chainNATOutput); delErr != nil {
|
||||||
|
log.Warnf("failed to rollback chain %s: %v", chainNATOutput, delErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fmt.Errorf("add OUTPUT jump rule: %w", err)
|
||||||
|
}
|
||||||
|
r.rules[jumpNATOutput] = jumpRule
|
||||||
|
|
||||||
|
r.updateState()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddOutputDNAT adds an OUTPUT chain DNAT rule for locally-generated traffic.
|
||||||
|
func (r *family) AddOutputDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
||||||
|
ruleID := firewall.RuleID(fmt.Sprintf("output-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
||||||
|
|
||||||
|
if _, exists := r.rules[ruleID]; exists {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.ensureNATOutputChain(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
dnatRule := []string{
|
||||||
|
"-p", strings.ToLower(protoForFamily(protocol, localAddr.Is6())),
|
||||||
|
"--dport", strconv.Itoa(int(originalPort)),
|
||||||
|
"-d", localAddr.String(),
|
||||||
|
"-j", "DNAT",
|
||||||
|
"--to-destination", ":" + strconv.Itoa(int(translatedPort)),
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.iptablesClient.Append(tableNat, chainNATOutput, dnatRule...); err != nil {
|
||||||
|
return fmt.Errorf("add output DNAT rule: %w", err)
|
||||||
|
}
|
||||||
|
r.rules[ruleID] = dnatRule
|
||||||
|
|
||||||
|
r.updateState()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
|
||||||
|
func (r *family) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
||||||
|
ruleID := firewall.RuleID(fmt.Sprintf("output-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
||||||
|
|
||||||
|
if dnatRule, exists := r.rules[ruleID]; exists {
|
||||||
|
if err := r.iptablesClient.Delete(tableNat, chainNATOutput, dnatRule...); err != nil {
|
||||||
|
return fmt.Errorf("delete output DNAT rule: %w", err)
|
||||||
|
}
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.updateState()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -80,7 +80,7 @@ func iptDnatV6(port uint16) fw.ForwardRule {
|
|||||||
// and a single DisableRouting drops both back to zero.
|
// and a single DisableRouting drops both back to zero.
|
||||||
func TestIptablesRouting_RepeatedEnableSingleReference(t *testing.T) {
|
func TestIptablesRouting_RepeatedEnableSingleReference(t *testing.T) {
|
||||||
m := newIptRefcountManager(t, true)
|
m := newIptRefcountManager(t, true)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
require.NoError(t, m.EnableRouting(), "first enable")
|
require.NoError(t, m.EnableRouting(), "first enable")
|
||||||
require.NoError(t, m.EnableRouting(), "second enable")
|
require.NoError(t, m.EnableRouting(), "second enable")
|
||||||
@@ -99,7 +99,7 @@ func TestIptablesRouting_RepeatedEnableSingleReference(t *testing.T) {
|
|||||||
// DisableRouting does not release references held by active DNAT rules.
|
// DisableRouting does not release references held by active DNAT rules.
|
||||||
func TestIptablesRouting_DisableKeepsDNATReference(t *testing.T) {
|
func TestIptablesRouting_DisableKeepsDNATReference(t *testing.T) {
|
||||||
m := newIptRefcountManager(t, true)
|
m := newIptRefcountManager(t, true)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(iptDnatV6(9095))
|
r1, err := m.AddDNATRule(iptDnatV6(9095))
|
||||||
require.NoError(t, err, "add v6 dnat")
|
require.NoError(t, err, "add v6 dnat")
|
||||||
@@ -116,7 +116,7 @@ func TestIptablesRouting_DisableKeepsDNATReference(t *testing.T) {
|
|||||||
// TestIptablesDNAT_RefcountBalancedV4 covers a Balanced Add/Delete pair on v4.
|
// TestIptablesDNAT_RefcountBalancedV4 covers a Balanced Add/Delete pair on v4.
|
||||||
func TestIptablesDNAT_RefcountBalancedV4(t *testing.T) {
|
func TestIptablesDNAT_RefcountBalancedV4(t *testing.T) {
|
||||||
m := newIptRefcountManager(t, false)
|
m := newIptRefcountManager(t, false)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(iptDnatV4(7081))
|
r1, err := m.AddDNATRule(iptDnatV4(7081))
|
||||||
require.NoError(t, err, "add v4 dnat 1")
|
require.NoError(t, err, "add v4 dnat 1")
|
||||||
@@ -145,9 +145,9 @@ func TestIptablesDNAT_RefcountBalancedV4(t *testing.T) {
|
|||||||
// decrements back to zero.
|
// decrements back to zero.
|
||||||
func TestIptablesDNAT_RefcountBalancedV6(t *testing.T) {
|
func TestIptablesDNAT_RefcountBalancedV6(t *testing.T) {
|
||||||
m := newIptRefcountManager(t, true)
|
m := newIptRefcountManager(t, true)
|
||||||
require.NotNil(t, m.router6, "v6 router")
|
require.NotNil(t, m.family6, "v6 family")
|
||||||
require.Same(t, m.router.ipFwdState, m.router6.ipFwdState, "shared state")
|
require.Same(t, m.family4.ipFwdState, m.family6.ipFwdState, "shared state")
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(iptDnatV6(9081))
|
r1, err := m.AddDNATRule(iptDnatV6(9081))
|
||||||
require.NoError(t, err, "add v6 dnat 1")
|
require.NoError(t, err, "add v6 dnat 1")
|
||||||
@@ -176,7 +176,7 @@ func TestIptablesDNAT_RefcountBalancedV6(t *testing.T) {
|
|||||||
// without bumping the refcount.
|
// without bumping the refcount.
|
||||||
func TestIptablesDNAT_DuplicateAddNoLeak(t *testing.T) {
|
func TestIptablesDNAT_DuplicateAddNoLeak(t *testing.T) {
|
||||||
m := newIptRefcountManager(t, true)
|
m := newIptRefcountManager(t, true)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
rule := iptDnatV4(7083)
|
rule := iptDnatV4(7083)
|
||||||
r1, err := m.AddDNATRule(rule)
|
r1, err := m.AddDNATRule(rule)
|
||||||
@@ -198,7 +198,7 @@ func TestIptablesDNAT_DuplicateAddNoLeak(t *testing.T) {
|
|||||||
// neither errors nor releases the refcount.
|
// neither errors nor releases the refcount.
|
||||||
func TestIptablesDNAT_DeleteMissingNoUnderflow(t *testing.T) {
|
func TestIptablesDNAT_DeleteMissingNoUnderflow(t *testing.T) {
|
||||||
m := newIptRefcountManager(t, true)
|
m := newIptRefcountManager(t, true)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
phantom := iptDnatV4(7099)
|
phantom := iptDnatV4(7099)
|
||||||
require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4")
|
require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4")
|
||||||
@@ -223,7 +223,7 @@ func TestIptablesDNAT_DeleteMissingNoUnderflow(t *testing.T) {
|
|||||||
// rule is a no-op.
|
// rule is a no-op.
|
||||||
func TestIptablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) {
|
func TestIptablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) {
|
||||||
m := newIptRefcountManager(t, true)
|
m := newIptRefcountManager(t, true)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(iptDnatV6(9083))
|
r1, err := m.AddDNATRule(iptDnatV6(9083))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|||||||
@@ -0,0 +1,260 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package iptables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"maps"
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/coreos/go-iptables/iptables"
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
nbid "github.com/netbirdio/netbird/client/internal/acl/id"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/routemanager/ipfwdstate"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/routemanager/refcounter"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/statemanager"
|
||||||
|
)
|
||||||
|
|
||||||
|
// constants needed to manage and create iptable rules
|
||||||
|
const (
|
||||||
|
tableFilter = "filter"
|
||||||
|
tableNat = "nat"
|
||||||
|
tableMangle = "mangle"
|
||||||
|
tableRaw = "raw"
|
||||||
|
|
||||||
|
// chainACLInput is the peer ACL chain that holds installed
|
||||||
|
// peer-filtering rules.
|
||||||
|
chainACLInput = "NETBIRD-ACL-INPUT"
|
||||||
|
|
||||||
|
// mangleForwardKey is the entries map key for mangle FORWARD guard
|
||||||
|
// rules that prevent external DNAT from bypassing ACL rules.
|
||||||
|
mangleForwardKey chainKey = "MANGLE-FORWARD"
|
||||||
|
|
||||||
|
chainInput = "INPUT"
|
||||||
|
chainOutput = "OUTPUT"
|
||||||
|
chainPostrouting = "POSTROUTING"
|
||||||
|
chainPrerouting = "PREROUTING"
|
||||||
|
chainForward = "FORWARD"
|
||||||
|
chainRTNAT = "NETBIRD-RT-NAT"
|
||||||
|
chainRTFwdIn = "NETBIRD-RT-FWD-IN"
|
||||||
|
chainRTFwdOut = "NETBIRD-RT-FWD-OUT"
|
||||||
|
chainRTPre = "NETBIRD-RT-PRE"
|
||||||
|
chainRTRdr = "NETBIRD-RT-RDR"
|
||||||
|
chainNATOutput = "NETBIRD-NAT-OUTPUT"
|
||||||
|
chainRTMSSClamp = "NETBIRD-RT-MSSCLAMP"
|
||||||
|
|
||||||
|
jumpManglePre = "jump-mangle-pre"
|
||||||
|
jumpNATPre = "jump-nat-pre"
|
||||||
|
jumpNATPost = "jump-nat-post"
|
||||||
|
jumpNATOutput = "jump-nat-output"
|
||||||
|
jumpMSSClamp = "jump-mss-clamp"
|
||||||
|
markManglePre = "mark-mangle-pre"
|
||||||
|
markManglePost = "mark-mangle-post"
|
||||||
|
matchSet = "--match-set"
|
||||||
|
|
||||||
|
dnatSuffix firewall.RuleID = "_dnat"
|
||||||
|
snatSuffix firewall.RuleID = "_snat"
|
||||||
|
fwdSuffix firewall.RuleID = "_fwd"
|
||||||
|
|
||||||
|
// ipv4TCPHeaderSize is the minimum IPv4 (20) + TCP (20) header size for MSS calculation.
|
||||||
|
ipv4TCPHeaderSize = 40
|
||||||
|
// ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation.
|
||||||
|
ipv6TCPHeaderSize = 60
|
||||||
|
)
|
||||||
|
|
||||||
|
type ruleInfo struct {
|
||||||
|
chain string
|
||||||
|
table string
|
||||||
|
rule []string
|
||||||
|
}
|
||||||
|
|
||||||
|
type routeRules map[firewall.RuleID][]string
|
||||||
|
|
||||||
|
// ruleSpec is a single iptables rule expressed as its argument list
|
||||||
|
// (e.g. {"-i", "wg0", "-j", "DROP"}).
|
||||||
|
type ruleSpec []string
|
||||||
|
|
||||||
|
// chainKey identifies the chain a seeded entry belongs to. It holds
|
||||||
|
// built-in chain names ("INPUT", "FORWARD", "PREROUTING") plus the
|
||||||
|
// synthetic mangleForwardKey bucket for the mangle FORWARD guard rules.
|
||||||
|
type chainKey string
|
||||||
|
|
||||||
|
// aclEntries maps a chain to the rules seeded into it to jump into or
|
||||||
|
// guard the netbird ACL chains.
|
||||||
|
type aclEntries map[chainKey][]ruleSpec
|
||||||
|
|
||||||
|
type entry struct {
|
||||||
|
spec ruleSpec
|
||||||
|
position int
|
||||||
|
}
|
||||||
|
|
||||||
|
// ipsetCounter is the shared hash:net refcounter used by peer and
|
||||||
|
// route ACLs alike. The ipset library does not support comments, so
|
||||||
|
// the key is just the set name (string).
|
||||||
|
type ipsetCounter = refcounter.Counter[string, []netip.Prefix, struct{}]
|
||||||
|
|
||||||
|
// family holds the per-address-family iptables state. One instance
|
||||||
|
// handles route ACLs, peer ACLs, NAT, DNAT, and MSS clamping for a
|
||||||
|
// single family; the top-level Manager owns one for v4 and another
|
||||||
|
// for v6.
|
||||||
|
type family struct {
|
||||||
|
iptablesClient *iptables.IPTables
|
||||||
|
wgIface iFaceMapper
|
||||||
|
v6 bool
|
||||||
|
|
||||||
|
// Peer ACL chain bookkeeping.
|
||||||
|
entries aclEntries
|
||||||
|
optionalEntries map[chainKey][]entry
|
||||||
|
|
||||||
|
// filters holds peer + route filter rules keyed by content hash.
|
||||||
|
// AddFilterRule writes here; DeleteFilterRule looks up by id.
|
||||||
|
filters map[nbid.RuleID]*Rule
|
||||||
|
ipsetCounter *ipsetCounter
|
||||||
|
// ipsetSupported records whether the kernel can create the hash:net
|
||||||
|
// sets the source matches rely on; probed once at init. When false,
|
||||||
|
// multi-source rules expand to one rule per source prefix.
|
||||||
|
ipsetSupported bool
|
||||||
|
|
||||||
|
// rules holds NAT, jump, and MSS-clamping rules (auxiliary
|
||||||
|
// plumbing that isn't a filter rule).
|
||||||
|
rules routeRules
|
||||||
|
|
||||||
|
// Routing / NAT.
|
||||||
|
legacyManagement bool
|
||||||
|
mtu uint16
|
||||||
|
ipFwdState *ipfwdstate.IPForwardingState
|
||||||
|
|
||||||
|
stateManager *statemanager.Manager
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFamily(iptablesClient *iptables.IPTables, wgIface iFaceMapper, mtu uint16) (*family, error) {
|
||||||
|
r := &family{
|
||||||
|
iptablesClient: iptablesClient,
|
||||||
|
wgIface: wgIface,
|
||||||
|
v6: iptablesClient.Proto() == iptables.ProtocolIPv6,
|
||||||
|
entries: make(aclEntries),
|
||||||
|
optionalEntries: make(map[chainKey][]entry),
|
||||||
|
filters: make(map[nbid.RuleID]*Rule),
|
||||||
|
rules: make(routeRules),
|
||||||
|
mtu: mtu,
|
||||||
|
ipFwdState: ipfwdstate.NewIPForwardingState(wgIface.Name()),
|
||||||
|
}
|
||||||
|
|
||||||
|
r.ipsetCounter = refcounter.New(
|
||||||
|
func(name string, sources []netip.Prefix) (struct{}, error) {
|
||||||
|
return struct{}{}, r.createIpSet(name, sources)
|
||||||
|
},
|
||||||
|
func(name string, _ struct{}) error {
|
||||||
|
return r.deleteIpSet(name)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
return r, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// init wires the family to the state manager and installs both the
|
||||||
|
// route ACL containers and the peer ACL chain skeleton.
|
||||||
|
func (r *family) init(stateManager *statemanager.Manager) error {
|
||||||
|
r.stateManager = stateManager
|
||||||
|
|
||||||
|
r.ipsetSupported = r.probeIPSetSupport()
|
||||||
|
|
||||||
|
if err := r.cleanUpDefaultForwardRules(); err != nil {
|
||||||
|
log.Errorf("failed to clean up rules from FORWARD chain: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.createContainers(); err != nil {
|
||||||
|
return fmt.Errorf("create containers: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.setupDataPlaneMark(); err != nil {
|
||||||
|
log.Errorf("failed to set up data plane mark: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.seedInitialEntries()
|
||||||
|
r.seedInitialOptionalEntries()
|
||||||
|
|
||||||
|
if err := r.cleanAclChains(); err != nil {
|
||||||
|
return fmt.Errorf("clean acl chains: %w", err)
|
||||||
|
}
|
||||||
|
if err := r.createDefaultChains(); err != nil {
|
||||||
|
return fmt.Errorf("create default chains: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.updateState()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reset tears down all firewall state owned by this family. ACL
|
||||||
|
// chain cleanup runs before route-chain cleanup because the route
|
||||||
|
// chains are still referenced by FORWARD jumps installed during
|
||||||
|
// seedInitialEntries; deleting them first would trip EBUSY.
|
||||||
|
func (r *family) Reset() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
|
||||||
|
if err := r.cleanAclChains(); err != nil {
|
||||||
|
merr = multierror.Append(merr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.cleanUpDefaultForwardRules(); err != nil {
|
||||||
|
merr = multierror.Append(merr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.ipsetCounter.Flush(); err != nil {
|
||||||
|
merr = multierror.Append(merr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.cleanupDataPlaneMark(); err != nil {
|
||||||
|
merr = multierror.Append(merr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
clear(r.rules)
|
||||||
|
clear(r.filters)
|
||||||
|
r.updateState()
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) updateState() {
|
||||||
|
if r.stateManager == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var currentState *ShutdownState
|
||||||
|
if existing := r.stateManager.GetState(currentState); existing != nil {
|
||||||
|
if existingState, ok := existing.(*ShutdownState); ok {
|
||||||
|
currentState = existingState
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if currentState == nil {
|
||||||
|
currentState = &ShutdownState{}
|
||||||
|
}
|
||||||
|
|
||||||
|
currentState.Lock()
|
||||||
|
defer currentState.Unlock()
|
||||||
|
|
||||||
|
// Clone the rule maps so the persisted state holds a private snapshot.
|
||||||
|
// The live maps keep being mutated by subsequent rule operations while
|
||||||
|
// the state manager marshals the state from its periodic-save goroutine.
|
||||||
|
// Sharing the maps by reference races the two and aborts the process with
|
||||||
|
// a concurrent map iteration and write. The ipset counter guards itself
|
||||||
|
// during marshaling, so it can be shared directly.
|
||||||
|
if r.v6 {
|
||||||
|
currentState.RouteRules6 = maps.Clone(r.rules)
|
||||||
|
currentState.RouteIPsetCounter6 = r.ipsetCounter
|
||||||
|
currentState.ACLEntries6 = maps.Clone(r.entries)
|
||||||
|
} else {
|
||||||
|
currentState.RouteRules = maps.Clone(r.rules)
|
||||||
|
currentState.RouteIPsetCounter = r.ipsetCounter
|
||||||
|
currentState.ACLEntries = maps.Clone(r.entries)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.stateManager.UpdateState(currentState); err != nil {
|
||||||
|
log.Errorf("failed to update state: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,430 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package iptables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
nbid "github.com/netbirdio/netbird/client/internal/acl/id"
|
||||||
|
nbnet "github.com/netbirdio/netbird/client/net"
|
||||||
|
)
|
||||||
|
|
||||||
|
// AddFilterRule installs a packet-filtering rule. With destination
|
||||||
|
// empty, the rule goes to the peer ACL input chain plus a paired
|
||||||
|
// mangle PREROUTING rule for the redirect mark. With destination set
|
||||||
|
// (prefix or named set), it goes to the route ACL forward chain.
|
||||||
|
// Multi-source rules collapse to one iptables rule via the shared
|
||||||
|
// hash:net ipset.
|
||||||
|
func (r *family) AddFilterRule(
|
||||||
|
id []byte,
|
||||||
|
sources []netip.Prefix,
|
||||||
|
destination firewall.Network,
|
||||||
|
proto firewall.Protocol,
|
||||||
|
sPort *firewall.Port,
|
||||||
|
dPort *firewall.Port,
|
||||||
|
action firewall.Action,
|
||||||
|
) (firewall.Rule, error) {
|
||||||
|
ruleID := nbid.GenerateRuleID(sources, destination, proto, sPort, dPort, action)
|
||||||
|
if existing, ok := r.filters[ruleID]; ok {
|
||||||
|
return existing, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
rule, err := r.installFilterRules(ruleID, sources, destination, proto, sPort, dPort, action, r.ipsetSupported)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
r.filters[ruleID] = rule
|
||||||
|
r.updateState()
|
||||||
|
return rule, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// installFilterRules resolves the source matches and installs one
|
||||||
|
// iptables rule per match. It is more than one rule only when useIPSet
|
||||||
|
// is false and a multi-source rule has to be expanded per prefix.
|
||||||
|
func (r *family) installFilterRules(
|
||||||
|
ruleID nbid.RuleID,
|
||||||
|
sources []netip.Prefix,
|
||||||
|
destination firewall.Network,
|
||||||
|
proto firewall.Protocol,
|
||||||
|
sPort *firewall.Port,
|
||||||
|
dPort *firewall.Port,
|
||||||
|
action firewall.Action,
|
||||||
|
useIPSet bool,
|
||||||
|
) (*Rule, error) {
|
||||||
|
srcMatches, err := r.applySourceMatches(sources, useIPSet)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("apply source match: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rule, err := r.installFilterRule(ruleID, srcMatches, destination, proto, sPort, dPort, action)
|
||||||
|
if err != nil {
|
||||||
|
for _, srcMatch := range srcMatches {
|
||||||
|
r.dropSourceMatch(srcMatch)
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return rule, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) hasRule(id nbid.RuleID) bool {
|
||||||
|
_, ok := r.filters[id]
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// hasDNATRule reports whether this family owns the DNAT rule set for
|
||||||
|
// the given user id. DNAT rules live in r.rules under the well-known
|
||||||
|
// "<id>_dnat" key; the lookup here is used by Manager.DeleteDNATRule
|
||||||
|
// to pick the right family.
|
||||||
|
func (r *family) hasDNATRule(id firewall.RuleID) bool {
|
||||||
|
_, ok := r.rules[id+dnatSuffix]
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteFilterRule removes a previously installed filter rule. The
|
||||||
|
// rule's stored chain/table identify where to delete from; source set
|
||||||
|
// references are recovered from the spec via findSets and dropped
|
||||||
|
// from the shared ipset counter.
|
||||||
|
func (r *family) DeleteFilterRule(rule firewall.Rule) error {
|
||||||
|
ruleID := rule.ID()
|
||||||
|
pr, ok := r.filters[ruleID]
|
||||||
|
if !ok {
|
||||||
|
log.Debugf("filter rule %s not found", ruleID)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteIfExists keeps the deletes idempotent so a retry after a
|
||||||
|
// partial failure does not error on the parts already removed.
|
||||||
|
var merr *multierror.Error
|
||||||
|
for _, fs := range pr.allSpecs() {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableFilter, pr.chain, fs.specs...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete rule from %s: %w", pr.chain, err))
|
||||||
|
}
|
||||||
|
if fs.mangleSpecs != nil {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableMangle, chainRTPre, fs.mangleSpecs...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete mangle rule: %w", err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if merr != nil {
|
||||||
|
// Leave the rule tracked so the caller retries the remaining part.
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The rule is gone from iptables, so untrack it regardless of how the
|
||||||
|
// refcount decrement goes, but surface decrement failures so callers
|
||||||
|
// see the ipset desync. Only the primary spec can reference sets: the
|
||||||
|
// per-prefix expansion never uses them.
|
||||||
|
delete(r.filters, ruleID)
|
||||||
|
r.updateState()
|
||||||
|
if err := r.decrementSetCounter(pr.specs); err != nil {
|
||||||
|
return fmt.Errorf("drop source set references: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// findSets scans an iptables rule spec for "-m set --match-set <name>
|
||||||
|
// <dir>" fragments and returns the named sets in occurrence order.
|
||||||
|
// Used at delete time to drop ipsetCounter references.
|
||||||
|
func findSets(rule []string) []string {
|
||||||
|
var sets []string
|
||||||
|
for i, arg := range rule {
|
||||||
|
if arg == "-m" && i+3 < len(rule) && rule[i+1] == "set" && rule[i+2] == matchSet {
|
||||||
|
sets = append(sets, rule[i+3])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return sets
|
||||||
|
}
|
||||||
|
|
||||||
|
// sourceNetwork classifies a source-prefix list into the firewall.Network
|
||||||
|
// shape the rest of the spec-builder consumes: empty for match-any, a
|
||||||
|
// single prefix inline, or an ipset for multiple sources.
|
||||||
|
func sourceNetwork(sources []netip.Prefix) firewall.Network {
|
||||||
|
switch {
|
||||||
|
case len(sources) == 0:
|
||||||
|
return firewall.Network{}
|
||||||
|
case len(sources) == 1 && sources[0].Bits() == 0:
|
||||||
|
return firewall.Network{}
|
||||||
|
case len(sources) == 1:
|
||||||
|
return firewall.Network{Prefix: sources[0]}
|
||||||
|
default:
|
||||||
|
return firewall.Network{Set: firewall.NewPrefixSet(sources)}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// applySourceMatches returns one source match fragment per iptables
|
||||||
|
// rule needed for the sources: normally a single fragment (a set match,
|
||||||
|
// a direct -s match, or nil for match-any), and one -s fragment per
|
||||||
|
// prefix when a multi-source rule cannot use ipset. Per-prefix rules
|
||||||
|
// are the only form a kernel without the ipset modules can express.
|
||||||
|
func (r *family) applySourceMatches(sources []netip.Prefix, useIPSet bool) ([][]string, error) {
|
||||||
|
network := sourceNetwork(sources)
|
||||||
|
if !network.IsSet() || useIPSet {
|
||||||
|
match, err := r.applySourceMatch(network, sources)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return [][]string{match}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
matches := make([][]string, 0, len(sources))
|
||||||
|
for _, source := range sources {
|
||||||
|
matches = append(matches, []string{"-s", source.String()})
|
||||||
|
}
|
||||||
|
return matches, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// applySourceMatch returns the iptables match fragment for the rule's
|
||||||
|
// source. For a Set it increments the shared ipset's refcount; for a
|
||||||
|
// Prefix it emits a direct -s match; for the wildcard it returns nil.
|
||||||
|
func (r *family) applySourceMatch(network firewall.Network, prefixes []netip.Prefix) ([]string, error) {
|
||||||
|
switch {
|
||||||
|
case network.IsSet():
|
||||||
|
if r.ipsetCounter == nil {
|
||||||
|
return nil, fmt.Errorf("multi-source peer rule requires shared ipset counter")
|
||||||
|
}
|
||||||
|
name := r.ipsetName(network.Set.HashedName())
|
||||||
|
if _, err := r.ipsetCounter.Increment(name, prefixes); err != nil {
|
||||||
|
return nil, fmt.Errorf("ipset increment %s: %w", name, err)
|
||||||
|
}
|
||||||
|
return []string{"-m", "set", matchSet, name, "src"}, nil
|
||||||
|
case network.IsPrefix():
|
||||||
|
return []string{"-s", network.Prefix.String()}, nil
|
||||||
|
default:
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// dropSourceMatch undoes whatever applySourceMatch reserved when
|
||||||
|
// installing a rule fails. Safe to call when the spec is empty or holds
|
||||||
|
// only inline matchers. Decrement errors are logged but not returned:
|
||||||
|
// the install error is what the caller needs to see.
|
||||||
|
func (r *family) dropSourceMatch(srcMatch []string) {
|
||||||
|
if r.ipsetCounter == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
for _, name := range findSets(srcMatch) {
|
||||||
|
if _, err := r.ipsetCounter.Decrement(name); err != nil {
|
||||||
|
log.Errorf("rollback ipset decrement %s: %v", name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// decrementSetCounter drops ipset references owned by a raw rule spec
|
||||||
|
// stored in r.rules (NAT / legacy route entries). It returns an error
|
||||||
|
// aggregate so the caller surfaces decrement failures.
|
||||||
|
func (r *family) decrementSetCounter(rule []string) error {
|
||||||
|
if r.ipsetCounter == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var merr *multierror.Error
|
||||||
|
for _, name := range findSets(rule) {
|
||||||
|
if _, err := r.ipsetCounter.Decrement(name); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("decrement counter: %w", err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// installFilterRule assembles and writes the iptables filter-chain
|
||||||
|
// rules for one filter rule, one per source match fragment. With
|
||||||
|
// destination empty the rules land in the peer ACL input chain and each
|
||||||
|
// gets a paired mangle PREROUTING rule for the redirect mark. With
|
||||||
|
// destination set the rules land in the route ACL forward chain and
|
||||||
|
// there is no mangle pairing.
|
||||||
|
func (r *family) installFilterRule(
|
||||||
|
ruleID nbid.RuleID,
|
||||||
|
srcMatches [][]string,
|
||||||
|
destination firewall.Network,
|
||||||
|
protocol firewall.Protocol,
|
||||||
|
sPort, dPort *firewall.Port,
|
||||||
|
action firewall.Action,
|
||||||
|
) (*Rule, error) {
|
||||||
|
isRoute := !destination.IsZero()
|
||||||
|
|
||||||
|
proto := protoForFamily(protocol, r.v6)
|
||||||
|
|
||||||
|
var destExp []string
|
||||||
|
if isRoute {
|
||||||
|
var err error
|
||||||
|
destExp, err = r.applyNetwork("-d", destination, nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("apply network -d: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
matchSpecs := filterMatchSpecs(proto, sPort, dPort)
|
||||||
|
|
||||||
|
chain := chainACLInput
|
||||||
|
if isRoute {
|
||||||
|
chain = chainRTFwdIn
|
||||||
|
}
|
||||||
|
|
||||||
|
var installed []filterSpecs
|
||||||
|
for _, srcMatch := range srcMatches {
|
||||||
|
specs := slices.Clone(srcMatch)
|
||||||
|
specs = append(specs, destExp...)
|
||||||
|
specs = append(specs, matchSpecs...)
|
||||||
|
|
||||||
|
var mangleSpecs []string
|
||||||
|
if !isRoute {
|
||||||
|
mangleSpecs = slices.Clone(specs)
|
||||||
|
mangleSpecs = append(mangleSpecs,
|
||||||
|
"-i", r.wgIface.Name(),
|
||||||
|
"-m", "addrtype", "--dst-type", "LOCAL",
|
||||||
|
"-j", "MARK", "--set-xmark", fmt.Sprintf("%#x", nbnet.PreroutingFwmarkRedirected),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
specs = append(specs, "-j", actionToStr(action))
|
||||||
|
|
||||||
|
if err := r.insertFilterRule(chain, action, specs); err != nil {
|
||||||
|
// Leave nothing half-installed: the caller sees an error, so a
|
||||||
|
// partial rule would silently keep matching without being tracked.
|
||||||
|
r.removeFilterSpecs(chain, installed)
|
||||||
|
r.dropSourceMatch(destExp)
|
||||||
|
return nil, fmt.Errorf("install filter rule on %s: %w", chain, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The mangle redirect-mark rule is best effort: the filter rule itself
|
||||||
|
// is what enforces the ACL, so a mangle failure must not undo it. Drop
|
||||||
|
// the spec so teardown does not try to remove a rule that was not added.
|
||||||
|
if mangleSpecs != nil {
|
||||||
|
if err := r.iptablesClient.Append(tableMangle, chainRTPre, mangleSpecs...); err != nil {
|
||||||
|
log.Errorf("add mangle rule: %v", err)
|
||||||
|
mangleSpecs = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
installed = append(installed, filterSpecs{specs: specs, mangleSpecs: mangleSpecs})
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Rule{
|
||||||
|
id: ruleID,
|
||||||
|
specs: installed[0].specs,
|
||||||
|
mangleSpecs: installed[0].mangleSpecs,
|
||||||
|
extraRules: installed[1:],
|
||||||
|
chain: chain,
|
||||||
|
v6: r.v6,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// insertFilterRule writes one assembled rule spec into the given ACL
|
||||||
|
// chain. Peer ACL drops are inserted at position 1 so they precede the
|
||||||
|
// chain's catch-all; route ACL drops are inserted at position 2 to sit
|
||||||
|
// immediately after the established/related accept rule.
|
||||||
|
func (r *family) insertFilterRule(chain string, action firewall.Action, specs []string) error {
|
||||||
|
if action == firewall.ActionDrop {
|
||||||
|
pos := 1
|
||||||
|
if chain == chainRTFwdIn {
|
||||||
|
pos = 2
|
||||||
|
}
|
||||||
|
return r.iptablesClient.Insert(tableFilter, chain, pos, specs...)
|
||||||
|
}
|
||||||
|
return r.iptablesClient.Append(tableFilter, chain, specs...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// removeFilterSpecs deletes the already-installed rules of a partially
|
||||||
|
// applied filter rule.
|
||||||
|
func (r *family) removeFilterSpecs(chain string, installed []filterSpecs) {
|
||||||
|
for _, fs := range installed {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableFilter, chain, fs.specs...); err != nil {
|
||||||
|
log.Debugf("delete partial filter rule: %v", err)
|
||||||
|
}
|
||||||
|
if fs.mangleSpecs != nil {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableMangle, chainRTPre, fs.mangleSpecs...); err != nil {
|
||||||
|
log.Debugf("delete partial mangle rule: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// applyNetwork resolves a firewall.Network into the iptables match
|
||||||
|
// fragment for the given direction flag (-s or -d). Set networks
|
||||||
|
// increment the shared ipset refcount; prefixes emit a direct match;
|
||||||
|
// an empty network returns no spec ("match any").
|
||||||
|
func (r *family) applyNetwork(flag string, network firewall.Network, prefixes []netip.Prefix) ([]string, error) {
|
||||||
|
direction := "src"
|
||||||
|
if flag == "-d" {
|
||||||
|
direction = "dst"
|
||||||
|
}
|
||||||
|
|
||||||
|
if network.IsSet() {
|
||||||
|
// A destination set is populated later from DNS results, so unlike a
|
||||||
|
// source set it cannot be expanded into per-prefix rules. Without
|
||||||
|
// ipset such a rule is not expressible; report it instead of
|
||||||
|
// installing something broader than the policy allows.
|
||||||
|
if flag == "-d" && !r.ipsetSupported {
|
||||||
|
return nil, fmt.Errorf("destination set %s requires ipset (ip_set_hash_net and xt_set)", network.Set.HashedName())
|
||||||
|
}
|
||||||
|
|
||||||
|
name := r.ipsetName(network.Set.HashedName())
|
||||||
|
if _, err := r.ipsetCounter.Increment(name, prefixes); err != nil {
|
||||||
|
return nil, fmt.Errorf("create or get ipset: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return []string{"-m", "set", matchSet, name, direction}, nil
|
||||||
|
}
|
||||||
|
if network.IsPrefix() {
|
||||||
|
return []string{flag, network.Prefix.String()}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// nolint:nilnil
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// protoForFamily translates ICMP to ICMPv6 for ip6tables.
|
||||||
|
// ip6tables requires "ipv6-icmp" (or "icmpv6") instead of "icmp".
|
||||||
|
func protoForFamily(protocol firewall.Protocol, v6 bool) string {
|
||||||
|
if v6 && protocol == firewall.ProtocolICMP {
|
||||||
|
return "ipv6-icmp"
|
||||||
|
}
|
||||||
|
return string(protocol)
|
||||||
|
}
|
||||||
|
|
||||||
|
// filterMatchSpecs returns the proto/port match fragment for a
|
||||||
|
// filtering rule. The source match (-s or -m set) is built by the
|
||||||
|
// caller and prepended.
|
||||||
|
func filterMatchSpecs(protocol string, sPort, dPort *firewall.Port) (specs []string) {
|
||||||
|
if protocol != "all" {
|
||||||
|
specs = append(specs, "-p", protocol)
|
||||||
|
}
|
||||||
|
specs = append(specs, applyPort("--sport", sPort)...)
|
||||||
|
specs = append(specs, applyPort("--dport", dPort)...)
|
||||||
|
return specs
|
||||||
|
}
|
||||||
|
|
||||||
|
func actionToStr(action firewall.Action) string {
|
||||||
|
if action == firewall.ActionAccept {
|
||||||
|
return "ACCEPT"
|
||||||
|
}
|
||||||
|
return "DROP"
|
||||||
|
}
|
||||||
|
|
||||||
|
func applyPort(flag string, port *firewall.Port) []string {
|
||||||
|
if port == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if port.IsRange && len(port.Values) == 2 {
|
||||||
|
return []string{flag, fmt.Sprintf("%d:%d", port.Values[0], port.Values[1])}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(port.Values) > 1 {
|
||||||
|
portList := make([]string, len(port.Values))
|
||||||
|
for i, p := range port.Values {
|
||||||
|
portList[i] = strconv.Itoa(int(p))
|
||||||
|
}
|
||||||
|
return []string{"-m", "multiport", flag, strings.Join(portList, ",")}
|
||||||
|
}
|
||||||
|
|
||||||
|
return []string{flag, strconv.Itoa(int(port.Values[0]))}
|
||||||
|
}
|
||||||
@@ -0,0 +1,93 @@
|
|||||||
|
package iptables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/coreos/go-iptables/iptables"
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
)
|
||||||
|
|
||||||
|
// InterfaceAllower opens the NetBird interface on the iptables filter INPUT
|
||||||
|
// chain so the host firewall doesn't drop traffic the userspace firewall
|
||||||
|
// handles. It is the fallback used when nftables is unavailable (an
|
||||||
|
// iptables-legacy host).
|
||||||
|
//
|
||||||
|
// It opens INPUT only: the userspace router never forwards in the kernel.
|
||||||
|
// firewalld trust is handled by the uspfilter manager, not here.
|
||||||
|
type InterfaceAllower struct {
|
||||||
|
ifaceName string
|
||||||
|
ipt4 *iptables.IPTables
|
||||||
|
// ipt6 is nil when the interface has no IPv6 overlay address.
|
||||||
|
ipt6 *iptables.IPTables
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewInterfaceAllower builds an iptables allower for the interface. It returns
|
||||||
|
// an error when iptables is unavailable, so the caller can fall back to
|
||||||
|
// firewalld trust.
|
||||||
|
func NewInterfaceAllower(wgIface iFaceMapper) (*InterfaceAllower, error) {
|
||||||
|
ipt4, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("iptables not available: %w", err)
|
||||||
|
}
|
||||||
|
if _, err := ipt4.ListChains(tableFilter); err != nil {
|
||||||
|
return nil, fmt.Errorf("iptables filter table not available: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
a := &InterfaceAllower{ifaceName: wgIface.Name(), ipt4: ipt4}
|
||||||
|
|
||||||
|
// Missing v6 must not break the v4 path: open v4 only and continue.
|
||||||
|
if wgIface.Address().HasIPv6() {
|
||||||
|
ipt6, err := iptables.NewWithProtocol(iptables.ProtocolIPv6)
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("ip6tables not available, opening interface on v4 only: %v", err)
|
||||||
|
} else if _, err := ipt6.ListChains(tableFilter); err != nil {
|
||||||
|
log.Warnf("ip6tables filter table not available, opening interface on v4 only: %v", err)
|
||||||
|
} else {
|
||||||
|
a.ipt6 = ipt6
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return a, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply inserts the interface accept rule on the filter INPUT chain. It removes
|
||||||
|
// any stale rule first so an unclean exit (e.g. SIGKILL, where Close never ran)
|
||||||
|
// is recovered deterministically rather than accumulating duplicates.
|
||||||
|
func (a *InterfaceAllower) Apply() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
for _, ipt := range a.clients() {
|
||||||
|
if err := ipt.DeleteIfExists(tableFilter, chainInput, a.inputRule()...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("clean stale interface accept rule: %w", err))
|
||||||
|
}
|
||||||
|
if err := ipt.Insert(tableFilter, chainInput, 1, a.inputRule()...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("add interface accept rule: %w", err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close removes the interface accept rule.
|
||||||
|
func (a *InterfaceAllower) Close() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
for _, ipt := range a.clients() {
|
||||||
|
if err := ipt.DeleteIfExists(tableFilter, chainInput, a.inputRule()...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove interface accept rule: %w", err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *InterfaceAllower) inputRule() []string {
|
||||||
|
return []string{"-i", a.ifaceName, "-j", "ACCEPT"}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *InterfaceAllower) clients() []*iptables.IPTables {
|
||||||
|
clients := []*iptables.IPTables{a.ipt4}
|
||||||
|
if a.ipt6 != nil {
|
||||||
|
clients = append(clients, a.ipt6)
|
||||||
|
}
|
||||||
|
return clients
|
||||||
|
}
|
||||||
@@ -0,0 +1,131 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package iptables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
"github.com/lrh3321/ipset-go"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
)
|
||||||
|
|
||||||
|
// probeIPSetSupport checks whether the kernel can create the ipset type
|
||||||
|
// used for source and destination matches. On kernels lacking the
|
||||||
|
// required ipset hash module, set creation fails (e.g. "invalid
|
||||||
|
// argument"), which would otherwise fail every multi-source rule and
|
||||||
|
// leave traffic the policy permits blocked by the catch-all drop. When
|
||||||
|
// unsupported, multi-source rules fall back to one rule per prefix.
|
||||||
|
func (r *family) probeIPSetSupport() bool {
|
||||||
|
// Use a unique name so concurrent processes don't collide and we only ever
|
||||||
|
// destroy the set we created ourselves. ipset names are limited to 31 chars,
|
||||||
|
// so use a short random suffix.
|
||||||
|
probeName := "nb-probe-" + uuid.New().String()[:8]
|
||||||
|
|
||||||
|
if err := r.createIPSet(probeName); err != nil {
|
||||||
|
log.Warnf("ipset is not available (failed to create probe set: %v); "+
|
||||||
|
"falling back to per-IP iptables ACL rules. Ensure the kernel provides "+
|
||||||
|
"the ipset hash:net module (ip_set_hash_net) for better performance with large rule sets", err)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.destroyIPSet(probeName); err != nil {
|
||||||
|
log.Debugf("destroy ipset probe set %q: %v", probeName, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) createIpSet(setName string, sources []netip.Prefix) error {
|
||||||
|
if err := r.createIPSet(setName); err != nil {
|
||||||
|
return fmt.Errorf("create set %s: %w", setName, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, prefix := range sources {
|
||||||
|
if err := r.addPrefixToIPSet(setName, prefix); err != nil {
|
||||||
|
// The refcounter records nothing when this callback errors,
|
||||||
|
// so destroy the set or it leaks in the kernel. A partial
|
||||||
|
// source set would also fail-open for deny rules, so the
|
||||||
|
// rule must fail rather than install with a missing source.
|
||||||
|
if derr := r.destroyIPSet(setName); derr != nil {
|
||||||
|
log.Warnf("rollback ipset %s after add failure: %v", setName, derr)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("add element to set %s: %w", setName, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) deleteIpSet(setName string) error {
|
||||||
|
if err := r.destroyIPSet(setName); err != nil {
|
||||||
|
return fmt.Errorf("destroy set %s: %w", setName, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugf("deleted unused ipset %s", setName)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
|
||||||
|
name := r.ipsetName(set.HashedName())
|
||||||
|
var merr *multierror.Error
|
||||||
|
for _, prefix := range prefixes {
|
||||||
|
if err := r.addPrefixToIPSet(name, prefix); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("add prefix to ipset: %w", err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if merr == nil {
|
||||||
|
log.Debugf("updated set %s with prefixes %v", name, prefixes)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) ipsetName(name string) string {
|
||||||
|
if r.v6 {
|
||||||
|
return name + "-v6"
|
||||||
|
}
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) createIPSet(name string) error {
|
||||||
|
opts := ipset.CreateOptions{
|
||||||
|
Replace: true,
|
||||||
|
}
|
||||||
|
if r.v6 {
|
||||||
|
opts.Family = ipset.FamilyIPV6
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := ipset.Create(name, ipset.TypeHashNet, opts); err != nil {
|
||||||
|
return fmt.Errorf("create ipset %s: %w", name, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugf("created ipset %s with type hash:net", name)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addPrefixToIPSet(name string, prefix netip.Prefix) error {
|
||||||
|
addr := prefix.Addr()
|
||||||
|
ip := addr.AsSlice()
|
||||||
|
|
||||||
|
entry := &ipset.Entry{
|
||||||
|
IP: ip,
|
||||||
|
CIDR: uint8(prefix.Bits()),
|
||||||
|
Replace: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := ipset.Add(name, entry); err != nil {
|
||||||
|
return fmt.Errorf("add prefix to ipset %s: %w", name, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) destroyIPSet(name string) error {
|
||||||
|
return ipset.Destroy(name)
|
||||||
|
}
|
||||||
@@ -3,7 +3,6 @@ package iptables
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
@@ -18,25 +17,20 @@ import (
|
|||||||
"github.com/netbirdio/netbird/client/internal/statemanager"
|
"github.com/netbirdio/netbird/client/internal/statemanager"
|
||||||
)
|
)
|
||||||
|
|
||||||
type resetter interface {
|
// Manager of iptables firewall. Per-family state (peer ACLs, route
|
||||||
Reset() error
|
// ACLs, NAT, DNAT, MSS clamping) lives on family; Manager dispatches
|
||||||
}
|
// by family and provides the public firewall.Manager surface.
|
||||||
|
|
||||||
// Manager of iptables firewall
|
|
||||||
type Manager struct {
|
type Manager struct {
|
||||||
mutex sync.Mutex
|
mutex sync.Mutex
|
||||||
|
|
||||||
wgIface iFaceMapper
|
wgIface iFaceMapper
|
||||||
|
|
||||||
ipv4Client *iptables.IPTables
|
ipv4Client *iptables.IPTables
|
||||||
aclMgr *aclManager
|
family4 *family
|
||||||
router *router
|
|
||||||
rawSupported bool
|
|
||||||
|
|
||||||
// IPv6 counterparts, nil when no v6 overlay
|
// IPv6 counterparts, nil when no v6 overlay
|
||||||
ipv6Client *iptables.IPTables
|
ipv6Client *iptables.IPTables
|
||||||
aclMgr6 *aclManager
|
family6 *family
|
||||||
router6 *router
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// iFaceMapper defines subset methods of interface required for manager
|
// iFaceMapper defines subset methods of interface required for manager
|
||||||
@@ -57,14 +51,9 @@ func Create(wgIface iFaceMapper, mtu uint16) (*Manager, error) {
|
|||||||
ipv4Client: iptablesClient,
|
ipv4Client: iptablesClient,
|
||||||
}
|
}
|
||||||
|
|
||||||
m.router, err = newRouter(iptablesClient, wgIface, mtu)
|
m.family4, err = newFamily(iptablesClient, wgIface, mtu)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("create router: %w", err)
|
return nil, fmt.Errorf("create family: %w", err)
|
||||||
}
|
|
||||||
|
|
||||||
m.aclMgr, err = newAclManager(iptablesClient, wgIface)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("create acl manager: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if wgIface.Address().HasIPv6() {
|
if wgIface.Address().HasIPv6() {
|
||||||
@@ -81,21 +70,18 @@ func (m *Manager) createIPv6Components(wgIface iFaceMapper, mtu uint16) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("init ip6tables: %w", err)
|
return fmt.Errorf("init ip6tables: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
family6, err := newFamily(ip6Client, wgIface, mtu)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("create v6 family: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Share the same IP forwarding state with the v4 family, since the
|
||||||
|
// forwarding refcounter is per-family but shared between both families.
|
||||||
|
family6.ipFwdState = m.family4.ipFwdState
|
||||||
|
|
||||||
m.ipv6Client = ip6Client
|
m.ipv6Client = ip6Client
|
||||||
|
m.family6 = family6
|
||||||
m.router6, err = newRouter(ip6Client, wgIface, mtu)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("create v6 router: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Share the same IP forwarding state with the v4 router, since
|
|
||||||
// Forwarding refcounter is per-family but shared between v4 and v6 routers.
|
|
||||||
m.router6.ipFwdState = m.router.ipFwdState
|
|
||||||
|
|
||||||
m.aclMgr6, err = newAclManager(ip6Client, wgIface)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("create v6 acl manager: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -109,7 +95,7 @@ func (m *Manager) Init(stateManager *statemanager.Manager) error {
|
|||||||
InterfaceState: &InterfaceState{
|
InterfaceState: &InterfaceState{
|
||||||
NameStr: m.wgIface.Name(),
|
NameStr: m.wgIface.Name(),
|
||||||
WGAddress: m.wgIface.Address(),
|
WGAddress: m.wgIface.Address(),
|
||||||
MTU: m.router.mtu,
|
MTU: m.family4.mtu,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
stateManager.RegisterState(state)
|
stateManager.RegisterState(state)
|
||||||
@@ -121,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 {
|
||||||
@@ -141,31 +123,24 @@ func (m *Manager) Init(stateManager *statemanager.Manager) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// initChains initializes router and ACL chains for both address families,
|
// initChains initializes the per-family firewall state for both
|
||||||
// rolling back on failure.
|
// address families, rolling back on failure.
|
||||||
func (m *Manager) initChains(stateManager *statemanager.Manager) error {
|
func (m *Manager) initChains(stateManager *statemanager.Manager) error {
|
||||||
type initStep struct {
|
type initStep struct {
|
||||||
name string
|
name string
|
||||||
init func(*statemanager.Manager) error
|
r *family
|
||||||
mgr resetter
|
|
||||||
}
|
}
|
||||||
|
|
||||||
steps := []initStep{
|
steps := []initStep{{"v4", m.family4}}
|
||||||
{"router", m.router.init, m.router},
|
|
||||||
{"acl manager", m.aclMgr.init, m.aclMgr},
|
|
||||||
}
|
|
||||||
if m.hasIPv6() {
|
if m.hasIPv6() {
|
||||||
steps = append(steps,
|
steps = append(steps, initStep{"v6", m.family6})
|
||||||
initStep{"v6 router", m.router6.init, m.router6},
|
|
||||||
initStep{"v6 acl manager", m.aclMgr6.init, m.aclMgr6},
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var initialized []initStep
|
var initialized []initStep
|
||||||
for _, s := range steps {
|
for _, s := range steps {
|
||||||
if err := s.init(stateManager); err != nil {
|
if err := s.r.init(stateManager); err != nil {
|
||||||
for i := len(initialized) - 1; i >= 0; i-- {
|
for i := len(initialized) - 1; i >= 0; i-- {
|
||||||
if rerr := initialized[i].mgr.Reset(); rerr != nil {
|
if rerr := initialized[i].r.Reset(); rerr != nil {
|
||||||
log.Warnf("rollback %s: %v", initialized[i].name, rerr)
|
log.Warnf("rollback %s: %v", initialized[i].name, rerr)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -176,84 +151,50 @@ func (m *Manager) initChains(stateManager *statemanager.Manager) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddPeerFiltering adds a rule to the firewall
|
// AddFilterRule installs a packet-filtering rule. See firewall.Manager
|
||||||
//
|
// docs for destination semantics. Sources are a single address family;
|
||||||
// Comment will be ignored because some system this feature is not supported
|
// the rule is dispatched to the matching v4 / v6 backend.
|
||||||
func (m *Manager) AddPeerFiltering(
|
func (m *Manager) AddFilterRule(
|
||||||
id []byte,
|
|
||||||
ip net.IP,
|
|
||||||
proto firewall.Protocol,
|
|
||||||
sPort *firewall.Port,
|
|
||||||
dPort *firewall.Port,
|
|
||||||
action firewall.Action,
|
|
||||||
ipsetName string,
|
|
||||||
) ([]firewall.Rule, error) {
|
|
||||||
m.mutex.Lock()
|
|
||||||
defer m.mutex.Unlock()
|
|
||||||
|
|
||||||
if ip.To4() != nil {
|
|
||||||
return m.aclMgr.AddPeerFiltering(id, ip, proto, sPort, dPort, action, ipsetName)
|
|
||||||
}
|
|
||||||
if !m.hasIPv6() {
|
|
||||||
return nil, fmt.Errorf("add peer filtering for %s: %w", ip, firewall.ErrIPv6NotInitialized)
|
|
||||||
}
|
|
||||||
return m.aclMgr6.AddPeerFiltering(id, ip, proto, sPort, dPort, action, ipsetName)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Manager) AddRouteFiltering(
|
|
||||||
id []byte,
|
id []byte,
|
||||||
sources []netip.Prefix,
|
sources []netip.Prefix,
|
||||||
destination firewall.Network,
|
destination firewall.Network,
|
||||||
proto firewall.Protocol,
|
proto firewall.Protocol,
|
||||||
sPort, dPort *firewall.Port,
|
sPort *firewall.Port,
|
||||||
|
dPort *firewall.Port,
|
||||||
action firewall.Action,
|
action firewall.Action,
|
||||||
) (firewall.Rule, error) {
|
) (firewall.Rule, error) {
|
||||||
|
if len(sources) == 0 {
|
||||||
|
return nil, firewall.ErrNoSources
|
||||||
|
}
|
||||||
|
|
||||||
m.mutex.Lock()
|
m.mutex.Lock()
|
||||||
defer m.mutex.Unlock()
|
defer m.mutex.Unlock()
|
||||||
|
|
||||||
if isIPv6RouteRule(sources, destination) {
|
fam := m.family4
|
||||||
|
if isIPv6Rule(sources, destination) {
|
||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return nil, fmt.Errorf("add route filtering: %w", firewall.ErrIPv6NotInitialized)
|
return nil, fmt.Errorf("add filtering: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.AddRouteFiltering(id, sources, destination, proto, sPort, dPort, action)
|
fam = m.family6
|
||||||
}
|
}
|
||||||
|
return fam.AddFilterRule(id, sources, destination, proto, sPort, dPort, action)
|
||||||
return m.router.AddRouteFiltering(id, sources, destination, proto, sPort, dPort, action)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func isIPv6RouteRule(sources []netip.Prefix, destination firewall.Network) bool {
|
// DeleteFilterRule removes a rule previously added via AddFilterRule.
|
||||||
if destination.IsPrefix() {
|
// The rule is looked up by id in each family's filter cache.
|
||||||
return destination.Prefix.Addr().Is6()
|
func (m *Manager) DeleteFilterRule(rule firewall.Rule) error {
|
||||||
}
|
|
||||||
return len(sources) > 0 && sources[0].Addr().Is6()
|
|
||||||
}
|
|
||||||
|
|
||||||
// DeletePeerRule from the firewall by rule definition
|
|
||||||
func (m *Manager) DeletePeerRule(rule firewall.Rule) error {
|
|
||||||
m.mutex.Lock()
|
m.mutex.Lock()
|
||||||
defer m.mutex.Unlock()
|
defer m.mutex.Unlock()
|
||||||
|
|
||||||
if m.hasIPv6() && isIPv6IptRule(rule) {
|
id := rule.ID()
|
||||||
return m.aclMgr6.DeletePeerRule(rule)
|
if m.family4.hasRule(id) {
|
||||||
|
return m.family4.DeleteFilterRule(rule)
|
||||||
}
|
}
|
||||||
return m.aclMgr.DeletePeerRule(rule)
|
if m.hasIPv6() && m.family6.hasRule(id) {
|
||||||
}
|
return m.family6.DeleteFilterRule(rule)
|
||||||
|
|
||||||
func isIPv6IptRule(rule firewall.Rule) bool {
|
|
||||||
r, ok := rule.(*Rule)
|
|
||||||
return ok && r.v6
|
|
||||||
}
|
|
||||||
|
|
||||||
// DeleteRouteRule deletes a routing rule.
|
|
||||||
// Route rules are keyed by content hash. Check v4 first, try v6 if not found.
|
|
||||||
func (m *Manager) DeleteRouteRule(rule firewall.Rule) error {
|
|
||||||
m.mutex.Lock()
|
|
||||||
defer m.mutex.Unlock()
|
|
||||||
|
|
||||||
if m.hasIPv6() && !m.router.hasRule(rule.ID()) {
|
|
||||||
return m.router6.DeleteRouteRule(rule)
|
|
||||||
}
|
}
|
||||||
return m.router.DeleteRouteRule(rule)
|
log.Debugf("filter rule %s not found in any family", id)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *Manager) IsServerRouteSupported() bool {
|
func (m *Manager) IsServerRouteSupported() bool {
|
||||||
@@ -272,10 +213,10 @@ func (m *Manager) AddNatRule(pair firewall.RouterPair) error {
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return fmt.Errorf("add NAT rule: %w", firewall.ErrIPv6NotInitialized)
|
return fmt.Errorf("add NAT rule: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.AddNatRule(pair)
|
return m.family6.AddNatRule(pair)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := m.router.AddNatRule(pair); err != nil {
|
if err := m.family4.AddNatRule(pair); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -284,7 +225,7 @@ func (m *Manager) AddNatRule(pair firewall.RouterPair) error {
|
|||||||
// wildcard 0.0.0.0/0 destination where the client resolves DNS.
|
// wildcard 0.0.0.0/0 destination where the client resolves DNS.
|
||||||
if m.hasIPv6() && pair.Dynamic {
|
if m.hasIPv6() && pair.Dynamic {
|
||||||
v6Pair := firewall.ToV6NatPair(pair)
|
v6Pair := firewall.ToV6NatPair(pair)
|
||||||
if err := m.router6.AddNatRule(v6Pair); err != nil {
|
if err := m.family6.AddNatRule(v6Pair); err != nil {
|
||||||
return fmt.Errorf("add v6 NAT rule: %w", err)
|
return fmt.Errorf("add v6 NAT rule: %w", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -300,18 +241,18 @@ func (m *Manager) RemoveNatRule(pair firewall.RouterPair) error {
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return m.router6.RemoveNatRule(pair)
|
return m.family6.RemoveNatRule(pair)
|
||||||
}
|
}
|
||||||
|
|
||||||
var merr *multierror.Error
|
var merr *multierror.Error
|
||||||
|
|
||||||
if err := m.router.RemoveNatRule(pair); err != nil {
|
if err := m.family4.RemoveNatRule(pair); err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("remove v4 NAT rule: %w", err))
|
merr = multierror.Append(merr, fmt.Errorf("remove v4 NAT rule: %w", err))
|
||||||
}
|
}
|
||||||
|
|
||||||
if m.hasIPv6() && pair.Dynamic {
|
if m.hasIPv6() && pair.Dynamic {
|
||||||
v6Pair := firewall.ToV6NatPair(pair)
|
v6Pair := firewall.ToV6NatPair(pair)
|
||||||
if err := m.router6.RemoveNatRule(v6Pair); err != nil {
|
if err := m.family6.RemoveNatRule(v6Pair); err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("remove v6 NAT rule: %w", err))
|
merr = multierror.Append(merr, fmt.Errorf("remove v6 NAT rule: %w", err))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -320,11 +261,14 @@ func (m *Manager) RemoveNatRule(pair firewall.RouterPair) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (m *Manager) SetLegacyManagement(isLegacy bool) error {
|
func (m *Manager) SetLegacyManagement(isLegacy bool) error {
|
||||||
if err := firewall.SetLegacyManagement(m.router, isLegacy); err != nil {
|
m.mutex.Lock()
|
||||||
|
defer m.mutex.Unlock()
|
||||||
|
|
||||||
|
if err := firewall.SetLegacyManagement(m.family4, isLegacy); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if m.hasIPv6() {
|
if m.hasIPv6() {
|
||||||
return firewall.SetLegacyManagement(m.router6, isLegacy)
|
return firewall.SetLegacyManagement(m.family6, isLegacy)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -336,24 +280,14 @@ 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.aclMgr6.Reset(); err != nil {
|
if err := m.family6.Reset(); err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("reset v6 acl manager: %w", err))
|
merr = multierror.Append(merr, fmt.Errorf("reset v6 family: %w", err))
|
||||||
}
|
|
||||||
if err := m.router6.Reset(); err != nil {
|
|
||||||
merr = multierror.Append(merr, fmt.Errorf("reset v6 router: %w", err))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := m.aclMgr.Reset(); err != nil {
|
if err := m.family4.Reset(); err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("reset acl manager: %w", err))
|
merr = multierror.Append(merr, fmt.Errorf("reset family: %w", err))
|
||||||
}
|
|
||||||
if err := m.router.Reset(); err != nil {
|
|
||||||
merr = multierror.Append(merr, fmt.Errorf("reset router: %w", err))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Appending to merr intentionally blocks DeleteState below so ShutdownState
|
// Appending to merr intentionally blocks DeleteState below so ShutdownState
|
||||||
@@ -372,27 +306,6 @@ func (m *Manager) Close(stateManager *statemanager.Manager) error {
|
|||||||
return nberrors.FormatErrorOrNil(merr)
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
}
|
}
|
||||||
|
|
||||||
// AllowNetbird allows netbird interface traffic.
|
|
||||||
// This is called when USPFilter wraps the native firewall, adding blanket accept
|
|
||||||
// rules so that packet filtering is handled in userspace instead of by netfilter.
|
|
||||||
func (m *Manager) AllowNetbird() error {
|
|
||||||
var merr *multierror.Error
|
|
||||||
if _, err := m.AddPeerFiltering(nil, net.IP{0, 0, 0, 0}, firewall.ProtocolALL, nil, nil, firewall.ActionAccept, ""); err != nil {
|
|
||||||
merr = multierror.Append(merr, fmt.Errorf("allow netbird v4 interface traffic: %w", err))
|
|
||||||
}
|
|
||||||
if m.hasIPv6() {
|
|
||||||
if _, err := m.AddPeerFiltering(nil, net.IPv6zero, firewall.ProtocolALL, nil, nil, firewall.ActionAccept, ""); err != nil {
|
|
||||||
merr = multierror.Append(merr, fmt.Errorf("allow netbird v6 interface traffic: %w", err))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := firewalld.TrustInterface(m.wgIface.Name()); err != nil {
|
|
||||||
log.Warnf("failed to trust interface in firewalld: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nberrors.FormatErrorOrNil(merr)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Flush doesn't need to be implemented for this manager
|
// Flush doesn't need to be implemented for this manager
|
||||||
func (m *Manager) Flush() error { return nil }
|
func (m *Manager) Flush() error { return nil }
|
||||||
|
|
||||||
@@ -403,11 +316,11 @@ func (m *Manager) SetLogLevel(log.Level) {
|
|||||||
|
|
||||||
func (m *Manager) EnableRouting() error {
|
func (m *Manager) EnableRouting() error {
|
||||||
// v6 only when the overlay actually has v6.
|
// v6 only when the overlay actually has v6.
|
||||||
return m.router.ipFwdState.RequestRouting(m.router6 != nil)
|
return m.family4.ipFwdState.RequestRouting(m.hasIPv6())
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *Manager) DisableRouting() error {
|
func (m *Manager) DisableRouting() error {
|
||||||
return m.router.ipFwdState.ReleaseRouting()
|
return m.family4.ipFwdState.ReleaseRouting()
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddDNATRule adds a DNAT rule
|
// AddDNATRule adds a DNAT rule
|
||||||
@@ -419,9 +332,9 @@ func (m *Manager) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error)
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return nil, fmt.Errorf("add DNAT rule: %w", firewall.ErrIPv6NotInitialized)
|
return nil, fmt.Errorf("add DNAT rule: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.AddDNATRule(rule)
|
return m.family6.AddDNATRule(rule)
|
||||||
}
|
}
|
||||||
return m.router.AddDNATRule(rule)
|
return m.family4.AddDNATRule(rule)
|
||||||
}
|
}
|
||||||
|
|
||||||
// DeleteDNATRule deletes a DNAT rule
|
// DeleteDNATRule deletes a DNAT rule
|
||||||
@@ -429,10 +342,10 @@ func (m *Manager) DeleteDNATRule(rule firewall.Rule) error {
|
|||||||
m.mutex.Lock()
|
m.mutex.Lock()
|
||||||
defer m.mutex.Unlock()
|
defer m.mutex.Unlock()
|
||||||
|
|
||||||
if m.hasIPv6() && !m.router.hasRule(rule.ID()+dnatSuffix) {
|
if m.hasIPv6() && !m.family4.hasDNATRule(rule.ID()) {
|
||||||
return m.router6.DeleteDNATRule(rule)
|
return m.family6.DeleteDNATRule(rule)
|
||||||
}
|
}
|
||||||
return m.router.DeleteDNATRule(rule)
|
return m.family4.DeleteDNATRule(rule)
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpdateSet updates the set with the given prefixes
|
// UpdateSet updates the set with the given prefixes
|
||||||
@@ -449,12 +362,12 @@ func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := m.router.UpdateSet(set, v4Prefixes); err != nil {
|
if err := m.family4.UpdateSet(set, v4Prefixes); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
if m.hasIPv6() && len(v6Prefixes) > 0 {
|
if m.hasIPv6() && len(v6Prefixes) > 0 {
|
||||||
if err := m.router6.UpdateSet(set, v6Prefixes); err != nil {
|
if err := m.family6.UpdateSet(set, v6Prefixes); err != nil {
|
||||||
return fmt.Errorf("update v6 set: %w", err)
|
return fmt.Errorf("update v6 set: %w", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -471,9 +384,9 @@ func (m *Manager) AddInboundDNAT(localAddr netip.Addr, protocol firewall.Protoco
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return fmt.Errorf("add inbound DNAT: %w", firewall.ErrIPv6NotInitialized)
|
return fmt.Errorf("add inbound DNAT: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.AddInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family6.AddInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
return m.router.AddInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family4.AddInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
|
|
||||||
// RemoveInboundDNAT removes an inbound DNAT rule.
|
// RemoveInboundDNAT removes an inbound DNAT rule.
|
||||||
@@ -485,9 +398,9 @@ func (m *Manager) RemoveInboundDNAT(localAddr netip.Addr, protocol firewall.Prot
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return fmt.Errorf("remove inbound DNAT: %w", firewall.ErrIPv6NotInitialized)
|
return fmt.Errorf("remove inbound DNAT: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.RemoveInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family6.RemoveInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
return m.router.RemoveInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family4.RemoveInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddOutputDNAT adds an OUTPUT chain DNAT rule for locally-generated traffic.
|
// AddOutputDNAT adds an OUTPUT chain DNAT rule for locally-generated traffic.
|
||||||
@@ -499,9 +412,9 @@ func (m *Manager) AddOutputDNAT(localAddr netip.Addr, protocol firewall.Protocol
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return fmt.Errorf("add output DNAT: %w", firewall.ErrIPv6NotInitialized)
|
return fmt.Errorf("add output DNAT: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.AddOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family6.AddOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
return m.router.AddOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family4.AddOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
|
|
||||||
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
|
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
|
||||||
@@ -513,139 +426,21 @@ func (m *Manager) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Proto
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return fmt.Errorf("remove output DNAT: %w", firewall.ErrIPv6NotInitialized)
|
return fmt.Errorf("remove output DNAT: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family6.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
return m.router.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"}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// isIPv6Rule reports whether the rule belongs to the IPv6 family, from
|
||||||
|
// the destination prefix when set, otherwise from the (single-family)
|
||||||
|
// sources.
|
||||||
|
func isIPv6Rule(sources []netip.Prefix, destination firewall.Network) bool {
|
||||||
|
if destination.IsPrefix() {
|
||||||
|
return destination.Prefix.Addr().Is6()
|
||||||
|
}
|
||||||
|
return len(sources) > 0 && sources[0].Addr().Is6()
|
||||||
|
}
|
||||||
|
|||||||
@@ -5,16 +5,19 @@ package iptables
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/coreos/go-iptables/iptables"
|
"github.com/coreos/go-iptables/iptables"
|
||||||
|
"github.com/lrh3321/ipset-go"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
fw "github.com/netbirdio/netbird/client/firewall/manager"
|
fw "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
"github.com/netbirdio/netbird/client/iface"
|
"github.com/netbirdio/netbird/client/iface"
|
||||||
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
||||||
|
"github.com/netbirdio/netbird/shared/management/domain"
|
||||||
)
|
)
|
||||||
|
|
||||||
var ifaceMock = &iFaceMock{
|
var ifaceMock = &iFaceMock{
|
||||||
@@ -67,47 +70,37 @@ func TestIptablesManager(t *testing.T) {
|
|||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
var rule2 []fw.Rule
|
var rule2 fw.Rule
|
||||||
t.Run("add second rule", func(t *testing.T) {
|
t.Run("add second rule", func(t *testing.T) {
|
||||||
ip := netip.MustParseAddr("10.20.0.3")
|
ip := netip.MustParseAddr("10.20.0.3")
|
||||||
port := &fw.Port{
|
port := &fw.Port{
|
||||||
IsRange: true,
|
IsRange: true,
|
||||||
Values: []uint16{8043, 8046},
|
Values: []uint16{8043, 8046},
|
||||||
}
|
}
|
||||||
rule2, err = manager.AddPeerFiltering(nil, ip.AsSlice(), "tcp", port, nil, fw.ActionAccept, "")
|
rule2, err = manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, "tcp", port, nil, fw.ActionAccept)
|
||||||
require.NoError(t, err, "failed to add rule")
|
require.NoError(t, err, "failed to add rule")
|
||||||
|
|
||||||
for _, r := range rule2 {
|
rr := rule2.(*Rule)
|
||||||
rr := r.(*Rule)
|
checkRuleSpecs(t, ipv4Client, rr.chain, true, rr.specs...)
|
||||||
checkRuleSpecs(t, ipv4Client, rr.chain, true, rr.specs...)
|
|
||||||
}
|
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("delete second rule", func(t *testing.T) {
|
t.Run("delete second rule", func(t *testing.T) {
|
||||||
for _, r := range rule2 {
|
require.NoError(t, manager.DeleteFilterRule(rule2), "failed to delete rule")
|
||||||
err := manager.DeletePeerRule(r)
|
|
||||||
require.NoError(t, err, "failed to delete rule")
|
|
||||||
}
|
|
||||||
|
|
||||||
require.Empty(t, manager.aclMgr.ipsetStore.ipsets, "rulesets index after removed second rule must be empty")
|
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("reset check", func(t *testing.T) {
|
t.Run("reset check", func(t *testing.T) {
|
||||||
// add second rule
|
// add second rule
|
||||||
ip := netip.MustParseAddr("10.20.0.3")
|
ip := netip.MustParseAddr("10.20.0.3")
|
||||||
port := &fw.Port{Values: []uint16{5353}}
|
port := &fw.Port{Values: []uint16{5353}}
|
||||||
_, err = manager.AddPeerFiltering(nil, ip.AsSlice(), "udp", nil, port, fw.ActionAccept, "")
|
_, err = manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, "udp", nil, port, fw.ActionAccept)
|
||||||
require.NoError(t, err, "failed to add rule")
|
require.NoError(t, err, "failed to add rule")
|
||||||
|
|
||||||
err = manager.Close(nil)
|
err = manager.Close(nil)
|
||||||
require.NoError(t, err, "failed to reset")
|
require.NoError(t, err, "failed to reset")
|
||||||
|
|
||||||
ok, err := ipv4Client.ChainExists("filter", chainNameInputRules)
|
ok, err := ipv4Client.ChainExists("filter", chainACLInput)
|
||||||
require.NoError(t, err, "failed check chain exists")
|
require.NoError(t, err, "failed check chain exists")
|
||||||
|
require.Falsef(t, ok, "chain %q still exists after Close", chainACLInput)
|
||||||
if ok {
|
|
||||||
require.NoErrorf(t, err, "chain '%v' still exists after Close", chainNameInputRules)
|
|
||||||
}
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -128,15 +121,13 @@ func TestIptablesManagerDenyRules(t *testing.T) {
|
|||||||
ip := netip.MustParseAddr("10.20.0.3")
|
ip := netip.MustParseAddr("10.20.0.3")
|
||||||
port := &fw.Port{Values: []uint16{22}}
|
port := &fw.Port{Values: []uint16{22}}
|
||||||
|
|
||||||
rule, err := manager.AddPeerFiltering(nil, ip.AsSlice(), "tcp", nil, port, fw.ActionDrop, "deny-ssh")
|
rule, err := manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, "tcp", nil, port, fw.ActionDrop)
|
||||||
require.NoError(t, err, "failed to add deny rule")
|
require.NoError(t, err, "failed to add deny rule")
|
||||||
require.NotEmpty(t, rule, "deny rule should not be empty")
|
require.NotNil(t, rule, "deny rule should not be nil")
|
||||||
|
|
||||||
// Verify the rule was added by checking iptables
|
// Verify the rule was added by checking iptables
|
||||||
for _, r := range rule {
|
rr := rule.(*Rule)
|
||||||
rr := r.(*Rule)
|
checkRuleSpecs(t, ipv4Client, rr.chain, true, rr.specs...)
|
||||||
checkRuleSpecs(t, ipv4Client, rr.chain, true, rr.specs...)
|
|
||||||
}
|
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("deny rule precedence test", func(t *testing.T) {
|
t.Run("deny rule precedence test", func(t *testing.T) {
|
||||||
@@ -144,36 +135,40 @@ func TestIptablesManagerDenyRules(t *testing.T) {
|
|||||||
port := &fw.Port{Values: []uint16{80}}
|
port := &fw.Port{Values: []uint16{80}}
|
||||||
|
|
||||||
// Add accept rule first
|
// Add accept rule first
|
||||||
_, err := manager.AddPeerFiltering(nil, ip.AsSlice(), "tcp", nil, port, fw.ActionAccept, "accept-http")
|
_, err := manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, "tcp", nil, port, fw.ActionAccept)
|
||||||
require.NoError(t, err, "failed to add accept rule")
|
require.NoError(t, err, "failed to add accept rule")
|
||||||
|
|
||||||
// Add deny rule second for same IP/port - this should take precedence
|
// Add deny rule second for same IP/port - this should take precedence
|
||||||
_, err = manager.AddPeerFiltering(nil, ip.AsSlice(), "tcp", nil, port, fw.ActionDrop, "deny-http")
|
_, err = manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, "tcp", nil, port, fw.ActionDrop)
|
||||||
require.NoError(t, err, "failed to add deny rule")
|
require.NoError(t, err, "failed to add deny rule")
|
||||||
|
|
||||||
// Inspect the actual iptables rules to verify deny rule comes before accept rule
|
// Inspect the actual iptables rules to verify deny rule comes before accept rule
|
||||||
rules, err := ipv4Client.List("filter", chainNameInputRules)
|
rules, err := ipv4Client.List("filter", chainACLInput)
|
||||||
require.NoError(t, err, "failed to list iptables rules")
|
require.NoError(t, err, "failed to list iptables rules")
|
||||||
|
|
||||||
// Debug: print all rules
|
// Debug: print all rules
|
||||||
t.Logf("All iptables rules in chain %s:", chainNameInputRules)
|
t.Logf("All iptables rules in chain %s:", chainACLInput)
|
||||||
for i, rule := range rules {
|
for i, rule := range rules {
|
||||||
t.Logf(" [%d] %s", i, rule)
|
t.Logf(" [%d] %s", i, rule)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Single-source rules emit a direct `-s <ip>/32 ... --dport 80`
|
||||||
|
// match. Match on that shape instead of the legacy
|
||||||
|
// per-(action,port) ipset names ("deny-http"/"accept-http")
|
||||||
|
// that this test predates.
|
||||||
|
srcMatch := fmt.Sprintf("-s %s/32", ip)
|
||||||
var denyRuleIndex, acceptRuleIndex = -1, -1
|
var denyRuleIndex, acceptRuleIndex = -1, -1
|
||||||
for i, rule := range rules {
|
for i, rule := range rules {
|
||||||
if strings.Contains(rule, "DROP") {
|
if !strings.Contains(rule, srcMatch) || !strings.Contains(rule, "--dport 80") {
|
||||||
t.Logf("Found DROP rule at index %d: %s", i, rule)
|
continue
|
||||||
if strings.Contains(rule, "deny-http") && strings.Contains(rule, "80") {
|
|
||||||
denyRuleIndex = i
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if strings.Contains(rule, "ACCEPT") {
|
if strings.Contains(rule, "-j DROP") {
|
||||||
|
t.Logf("Found DROP rule at index %d: %s", i, rule)
|
||||||
|
denyRuleIndex = i
|
||||||
|
}
|
||||||
|
if strings.Contains(rule, "-j ACCEPT") {
|
||||||
t.Logf("Found ACCEPT rule at index %d: %s", i, rule)
|
t.Logf("Found ACCEPT rule at index %d: %s", i, rule)
|
||||||
if strings.Contains(rule, "accept-http") && strings.Contains(rule, "80") {
|
acceptRuleIndex = i
|
||||||
acceptRuleIndex = i
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -198,7 +193,6 @@ func TestIptablesManagerIPSet(t *testing.T) {
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
// just check on the local interface
|
|
||||||
manager, err := Create(mock, iface.DefaultMTU)
|
manager, err := Create(mock, iface.DefaultMTU)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.NoError(t, manager.Init(nil))
|
require.NoError(t, manager.Init(nil))
|
||||||
@@ -212,27 +206,39 @@ func TestIptablesManagerIPSet(t *testing.T) {
|
|||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
var rule2 []fw.Rule
|
var rule2 fw.Rule
|
||||||
t.Run("add second rule", func(t *testing.T) {
|
t.Run("single source uses direct -s match (no ipset)", func(t *testing.T) {
|
||||||
ip := netip.MustParseAddr("10.20.0.3")
|
ip := netip.MustParseAddr("10.20.0.3")
|
||||||
port := &fw.Port{
|
port := &fw.Port{
|
||||||
Values: []uint16{443},
|
Values: []uint16{443},
|
||||||
}
|
}
|
||||||
rule2, err = manager.AddPeerFiltering(nil, ip.AsSlice(), "tcp", port, nil, fw.ActionAccept, "default")
|
rule2, err = manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, "tcp", port, nil, fw.ActionAccept)
|
||||||
for _, r := range rule2 {
|
require.NoError(t, err, "failed to add rule")
|
||||||
require.NoError(t, err, "failed to add rule")
|
require.NotNil(t, rule2)
|
||||||
require.Equal(t, r.(*Rule).ipsetName, "default-sport", "ipset name must be set")
|
require.Contains(t, rule2.(*Rule).specs, "-s",
|
||||||
require.Equal(t, r.(*Rule).ip, "10.20.0.3", "ipset IP must be set")
|
"single-source rule should use direct -s match, not an ipset")
|
||||||
}
|
require.Empty(t, findSets(rule2.(*Rule).specs),
|
||||||
|
"single-source rule should not allocate a shared ipset")
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("delete second rule", func(t *testing.T) {
|
t.Run("delete single-source rule", func(t *testing.T) {
|
||||||
for _, r := range rule2 {
|
require.NoError(t, manager.DeleteFilterRule(rule2), "failed to delete rule")
|
||||||
err := manager.DeletePeerRule(r)
|
})
|
||||||
require.NoError(t, err, "failed to delete rule")
|
|
||||||
|
|
||||||
require.Empty(t, manager.aclMgr.ipsetStore.ipsets, "rulesets index after removed second rule must be empty")
|
t.Run("multi-source uses shared ipset", func(t *testing.T) {
|
||||||
|
sources := []netip.Prefix{
|
||||||
|
netip.PrefixFrom(netip.MustParseAddr("10.20.0.3"), 32),
|
||||||
|
netip.PrefixFrom(netip.MustParseAddr("10.20.0.4"), 32),
|
||||||
|
netip.PrefixFrom(netip.MustParseAddr("10.20.0.5"), 32),
|
||||||
}
|
}
|
||||||
|
port := &fw.Port{Values: []uint16{8080}}
|
||||||
|
multi, err := manager.AddFilterRule(nil, sources, fw.Network{}, "tcp", nil, port, fw.ActionAccept)
|
||||||
|
require.NoError(t, err, "failed to add multi-source rule")
|
||||||
|
require.NotNil(t, multi, "multi-source rule must produce one iptables rule")
|
||||||
|
sets := findSets(multi.(*Rule).specs)
|
||||||
|
require.Len(t, sets, 1, "multi-source rule must reference exactly one ipset")
|
||||||
|
|
||||||
|
require.NoError(t, manager.DeleteFilterRule(multi))
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("reset check", func(t *testing.T) {
|
t.Run("reset check", func(t *testing.T) {
|
||||||
@@ -241,9 +247,324 @@ func TestIptablesManagerIPSet(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestIptablesFilterIPSetFallback verifies that when the kernel lacks
|
||||||
|
// ipset support, a multi-source rule falls back to one iptables rule
|
||||||
|
// per source prefix instead of silently leaving the chain empty. See
|
||||||
|
// discussion #6125.
|
||||||
|
func TestIptablesFilterIPSetFallback(t *testing.T) {
|
||||||
|
ipv4Client, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
manager, err := Create(ifaceMock, iface.DefaultMTU)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, manager.Init(nil))
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
require.NoError(t, manager.Close(nil))
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Simulate a kernel without the ipset hash module.
|
||||||
|
manager.family4.ipsetSupported = false
|
||||||
|
|
||||||
|
sources := []netip.Prefix{
|
||||||
|
netip.MustParsePrefix("10.20.0.42/32"),
|
||||||
|
netip.MustParsePrefix("10.20.0.43/32"),
|
||||||
|
}
|
||||||
|
port := &fw.Port{Values: []uint16{22}}
|
||||||
|
|
||||||
|
rule, err := manager.AddFilterRule(nil, sources, fw.Network{}, "tcp", nil, port, fw.ActionAccept)
|
||||||
|
require.NoError(t, err, "AddFilterRule should succeed via fallback")
|
||||||
|
|
||||||
|
rr := rule.(*Rule)
|
||||||
|
all := rr.allSpecs()
|
||||||
|
require.Len(t, all, len(sources), "each source prefix needs its own rule")
|
||||||
|
for i, fs := range all {
|
||||||
|
joined := strings.Join(fs.specs, " ")
|
||||||
|
require.Contains(t, joined, "-s "+sources[i].String(), "fallback rule must match by source prefix")
|
||||||
|
require.NotContains(t, joined, matchSet, "fallback rule must not use ipset matching")
|
||||||
|
|
||||||
|
// The rule must actually be present in the ACL chain (not silently dropped).
|
||||||
|
checkRuleSpecs(t, ipv4Client, rr.chain, true, fs.specs...)
|
||||||
|
|
||||||
|
// Every expanded peer rule keeps its own redirect-mark pairing.
|
||||||
|
require.NotNil(t, fs.mangleSpecs, "peer rule must carry a mangle pairing")
|
||||||
|
checkTableRuleSpecs(t, ipv4Client, tableMangle, chainRTPre, true, fs.mangleSpecs...)
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, manager.DeleteFilterRule(rule), "failed to delete fallback rule")
|
||||||
|
for _, fs := range all {
|
||||||
|
checkRuleSpecs(t, ipv4Client, rr.chain, false, fs.specs...)
|
||||||
|
checkTableRuleSpecs(t, ipv4Client, tableMangle, chainRTPre, false, fs.mangleSpecs...)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestIptablesFilterDestinationSetRequiresIPSet documents that a dynamic
|
||||||
|
// (domain) destination cannot be expressed without ipset: its prefixes are only
|
||||||
|
// known after DNS resolution, so there is nothing to expand into per-prefix
|
||||||
|
// rules. The call must report that rather than install a broader rule than the
|
||||||
|
// policy allows.
|
||||||
|
func TestIptablesFilterDestinationSetRequiresIPSet(t *testing.T) {
|
||||||
|
manager, err := Create(ifaceMock, iface.DefaultMTU)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, manager.Init(nil))
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
require.NoError(t, manager.Close(nil))
|
||||||
|
}()
|
||||||
|
|
||||||
|
manager.family4.ipsetSupported = false
|
||||||
|
|
||||||
|
destination := fw.Network{Set: fw.NewDomainSet(domain.List{"example.com"})}
|
||||||
|
|
||||||
|
_, err = manager.AddFilterRule(nil, []netip.Prefix{netip.MustParsePrefix("172.16.0.0/16")},
|
||||||
|
destination, fw.ProtocolALL, nil, nil, fw.ActionAccept)
|
||||||
|
require.Error(t, err, "a domain destination is not expressible without ipset")
|
||||||
|
require.ErrorContains(t, err, "requires ipset")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestIptablesNatRuleDropsSourceSetOnDestinationFailure covers a marking rule
|
||||||
|
// whose source set is created but whose destination set is not: the source
|
||||||
|
// reference has to go back, or the set it created stays in the kernel with a
|
||||||
|
// count nothing will ever drop.
|
||||||
|
func TestIptablesNatRuleDropsSourceSetOnDestinationFailure(t *testing.T) {
|
||||||
|
manager, err := Create(ifaceMock, iface.DefaultMTU)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, manager.Init(nil))
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
require.NoError(t, manager.Close(nil))
|
||||||
|
}()
|
||||||
|
|
||||||
|
sourceSet := fw.NewPrefixSet([]netip.Prefix{
|
||||||
|
netip.MustParsePrefix("100.0.0.0/16"),
|
||||||
|
netip.MustParsePrefix("10.10.0.0/16"),
|
||||||
|
})
|
||||||
|
destSet := fw.NewDomainSet(domain.List{"example.org"})
|
||||||
|
|
||||||
|
// Poison the destination set's name so its hash:net creation fails after
|
||||||
|
// the source set has already been created.
|
||||||
|
poisoned := manager.family4.ipsetName(destSet.HashedName())
|
||||||
|
require.NoError(t, ipset.Create(poisoned, ipset.TypeHashIP, ipset.CreateOptions{}))
|
||||||
|
t.Cleanup(func() {
|
||||||
|
if err := ipset.Destroy(poisoned); err != nil {
|
||||||
|
t.Logf("destroy poisoned set %s: %v", poisoned, err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
pair := fw.RouterPair{
|
||||||
|
ID: "nat-source-set-test",
|
||||||
|
Source: fw.Network{Set: sourceSet},
|
||||||
|
Destination: fw.Network{Set: destSet},
|
||||||
|
Masquerade: true,
|
||||||
|
Dynamic: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
require.Error(t, manager.AddNatRule(pair), "the destination set must fail to be created")
|
||||||
|
|
||||||
|
_, ok := manager.family4.ipsetCounter.Get(manager.family4.ipsetName(sourceSet.HashedName()))
|
||||||
|
require.False(t, ok, "the source set reference must be released")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestIptablesNatRuleReAddKeepsSetReferences re-adds the same NAT rule the way
|
||||||
|
// a repeated network-map update does. The marking rule's set references must not
|
||||||
|
// grow, or RemoveNatRule can never drop the count to zero and the set stays in
|
||||||
|
// the kernel for the rest of the process lifetime.
|
||||||
|
func TestIptablesNatRuleReAddKeepsSetReferences(t *testing.T) {
|
||||||
|
manager, err := Create(ifaceMock, iface.DefaultMTU)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, manager.Init(nil))
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
require.NoError(t, manager.Close(nil))
|
||||||
|
}()
|
||||||
|
|
||||||
|
set := fw.NewDomainSet(domain.List{"example.com"})
|
||||||
|
pair := fw.RouterPair{
|
||||||
|
ID: "nat-reference-test",
|
||||||
|
Source: fw.Network{Prefix: netip.MustParsePrefix("100.0.0.0/16")},
|
||||||
|
Destination: fw.Network{Set: set},
|
||||||
|
Masquerade: true,
|
||||||
|
Dynamic: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, manager.AddNatRule(pair), "add nat rule")
|
||||||
|
name := manager.family4.ipsetName(set.HashedName())
|
||||||
|
first, ok := manager.family4.ipsetCounter.Get(name)
|
||||||
|
require.True(t, ok, "the marking rule must hold a reference to its set")
|
||||||
|
|
||||||
|
require.NoError(t, manager.AddNatRule(pair), "re-add nat rule")
|
||||||
|
second, ok := manager.family4.ipsetCounter.Get(name)
|
||||||
|
require.True(t, ok, "the set must still be referenced")
|
||||||
|
require.Equal(t, first.Count, second.Count, "re-adding the same rule must not add references")
|
||||||
|
|
||||||
|
require.NoError(t, manager.RemoveNatRule(pair), "remove nat rule")
|
||||||
|
_, ok = manager.family4.ipsetCounter.Get(name)
|
||||||
|
require.False(t, ok, "removing the rule must drop the last reference")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestIptablesRouteFilterIPSetFallback covers the route ACL side of the
|
||||||
|
// fallback: with a destination set, the expanded per-source rules land
|
||||||
|
// in the route forward chain and are all removed on delete.
|
||||||
|
func TestIptablesRouteFilterIPSetFallback(t *testing.T) {
|
||||||
|
ipv4Client, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
manager, err := Create(ifaceMock, iface.DefaultMTU)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, manager.Init(nil))
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
require.NoError(t, manager.Close(nil))
|
||||||
|
}()
|
||||||
|
|
||||||
|
manager.family4.ipsetSupported = false
|
||||||
|
|
||||||
|
sources := []netip.Prefix{
|
||||||
|
netip.MustParsePrefix("172.16.0.0/16"),
|
||||||
|
netip.MustParsePrefix("192.168.0.0/16"),
|
||||||
|
}
|
||||||
|
destination := fw.Network{Prefix: netip.MustParsePrefix("10.0.0.0/8")}
|
||||||
|
port := &fw.Port{Values: []uint16{443}}
|
||||||
|
|
||||||
|
rule, err := manager.AddFilterRule(nil, sources, destination, "tcp", nil, port, fw.ActionAccept)
|
||||||
|
require.NoError(t, err, "route ACL must install without ipset")
|
||||||
|
|
||||||
|
rr := rule.(*Rule)
|
||||||
|
require.Equal(t, chainRTFwdIn, rr.chain, "route rule must land in the forward chain")
|
||||||
|
|
||||||
|
all := rr.allSpecs()
|
||||||
|
require.Len(t, all, len(sources), "each source prefix needs its own rule")
|
||||||
|
for i, fs := range all {
|
||||||
|
joined := strings.Join(fs.specs, " ")
|
||||||
|
require.Contains(t, joined, "-s "+sources[i].String(), "fallback rule must match by source prefix")
|
||||||
|
require.NotContains(t, joined, matchSet, "fallback rule must not use ipset matching")
|
||||||
|
require.Nil(t, fs.mangleSpecs, "route rules have no mangle pairing")
|
||||||
|
|
||||||
|
checkRuleSpecs(t, ipv4Client, rr.chain, true, fs.specs...)
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, manager.DeleteFilterRule(rule), "failed to delete fallback rule")
|
||||||
|
for _, fs := range all {
|
||||||
|
checkRuleSpecs(t, ipv4Client, rr.chain, false, fs.specs...)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestIptablesCloseRemovesAllState exercises a spread of rule kinds and then
|
||||||
|
// asserts Close puts every table it touches back exactly as it found it. A
|
||||||
|
// leaked chain, jump, or ipset survives the daemon and nothing can remove it
|
||||||
|
// afterwards, since the tracking that knew about it is gone.
|
||||||
|
func TestIptablesCloseRemovesAllState(t *testing.T) {
|
||||||
|
ipv4Client, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
before := snapshotIptables(t, ipv4Client)
|
||||||
|
|
||||||
|
manager, err := Create(ifaceMock, iface.DefaultMTU)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, manager.Init(nil))
|
||||||
|
|
||||||
|
// A failed assertion below returns before the Close under test, which would
|
||||||
|
// leave this test's chains and sets in the kernel for the next one.
|
||||||
|
t.Cleanup(func() {
|
||||||
|
if err := manager.Close(nil); err != nil {
|
||||||
|
t.Logf("close after failure: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
sources := []netip.Prefix{
|
||||||
|
netip.MustParsePrefix("10.20.0.42/32"),
|
||||||
|
netip.MustParsePrefix("10.20.0.43/32"),
|
||||||
|
}
|
||||||
|
|
||||||
|
// A multi-source peer rule: shared ipset plus the mangle redirect pairing.
|
||||||
|
_, err = manager.AddFilterRule(nil, sources, fw.Network{}, "tcp",
|
||||||
|
nil, &fw.Port{Values: []uint16{22}}, fw.ActionAccept)
|
||||||
|
require.NoError(t, err, "add peer rule")
|
||||||
|
|
||||||
|
// A route rule with a dynamic destination: a second set, in the forward chain.
|
||||||
|
_, err = manager.AddFilterRule(nil, sources,
|
||||||
|
fw.Network{Set: fw.NewDomainSet(domain.List{"example.com"})},
|
||||||
|
fw.ProtocolALL, nil, nil, fw.ActionDrop)
|
||||||
|
require.NoError(t, err, "add route rule")
|
||||||
|
|
||||||
|
// NAT marking for a routed destination, both directions.
|
||||||
|
pair := fw.RouterPair{
|
||||||
|
ID: "cleanup-test",
|
||||||
|
Source: fw.Network{Prefix: netip.MustParsePrefix("100.0.0.0/16")},
|
||||||
|
Destination: fw.Network{Prefix: netip.MustParsePrefix("192.168.55.0/24")},
|
||||||
|
Masquerade: true,
|
||||||
|
}
|
||||||
|
require.NoError(t, manager.AddNatRule(pair), "add nat rule")
|
||||||
|
require.NoError(t, manager.EnableRouting(), "enable routing")
|
||||||
|
|
||||||
|
// A DNAT redirect, which also holds a forwarding reference.
|
||||||
|
dnat := fw.ForwardRule{
|
||||||
|
Protocol: fw.ProtocolTCP,
|
||||||
|
DestinationPort: fw.Port{Values: []uint16{8080}},
|
||||||
|
TranslatedAddress: netip.MustParseAddr("10.20.0.44"),
|
||||||
|
TranslatedPort: fw.Port{Values: []uint16{80}},
|
||||||
|
}
|
||||||
|
_, err = manager.AddDNATRule(dnat)
|
||||||
|
require.NoError(t, err, "add dnat rule")
|
||||||
|
|
||||||
|
require.NotEqual(t, before, snapshotIptables(t, ipv4Client), "the manager must have installed state")
|
||||||
|
|
||||||
|
// Everything above stays in place, so Close is what has to remove it.
|
||||||
|
require.NoError(t, manager.Close(nil), "close")
|
||||||
|
|
||||||
|
after := snapshotIptables(t, ipv4Client)
|
||||||
|
require.Equal(t, before.chains, after.chains, "Close must remove every chain it created")
|
||||||
|
require.Equal(t, before.rules, after.rules, "Close must remove every rule it created")
|
||||||
|
require.Equal(t, before.sets, after.sets, "Close must destroy every ipset it created")
|
||||||
|
}
|
||||||
|
|
||||||
|
// iptablesState is a snapshot of the tables the manager writes to, used to
|
||||||
|
// compare the kernel before and after a manager lifetime.
|
||||||
|
type iptablesState struct {
|
||||||
|
chains map[string][]string
|
||||||
|
rules map[string][]string
|
||||||
|
sets []string
|
||||||
|
}
|
||||||
|
|
||||||
|
func snapshotIptables(t *testing.T, client *iptables.IPTables) iptablesState {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
state := iptablesState{
|
||||||
|
chains: map[string][]string{},
|
||||||
|
rules: map[string][]string{},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, table := range []string{tableFilter, tableNat, tableMangle, tableRaw} {
|
||||||
|
chains, err := client.ListChains(table)
|
||||||
|
require.NoErrorf(t, err, "list chains in %s", table)
|
||||||
|
slices.Sort(chains)
|
||||||
|
state.chains[table] = chains
|
||||||
|
|
||||||
|
for _, chain := range chains {
|
||||||
|
rules, err := client.List(table, chain)
|
||||||
|
require.NoErrorf(t, err, "list rules in %s/%s", table, chain)
|
||||||
|
state.rules[table+"/"+chain] = rules
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
sets, err := ipset.ListAll()
|
||||||
|
require.NoError(t, err, "list ipsets")
|
||||||
|
for _, set := range sets {
|
||||||
|
state.sets = append(state.sets, set.SetName)
|
||||||
|
}
|
||||||
|
slices.Sort(state.sets)
|
||||||
|
|
||||||
|
return state
|
||||||
|
}
|
||||||
|
|
||||||
func checkRuleSpecs(t *testing.T, ipv4Client *iptables.IPTables, chainName string, mustExists bool, rulespec ...string) {
|
func checkRuleSpecs(t *testing.T, ipv4Client *iptables.IPTables, chainName string, mustExists bool, rulespec ...string) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
exists, err := ipv4Client.Exists("filter", chainName, rulespec...)
|
checkTableRuleSpecs(t, ipv4Client, tableFilter, chainName, mustExists, rulespec...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func checkTableRuleSpecs(t *testing.T, ipv4Client *iptables.IPTables, table, chainName string, mustExists bool, rulespec ...string) {
|
||||||
|
t.Helper()
|
||||||
|
exists, err := ipv4Client.Exists(table, chainName, rulespec...)
|
||||||
require.NoError(t, err, "failed to check rule")
|
require.NoError(t, err, "failed to check rule")
|
||||||
require.Falsef(t, !exists && mustExists, "rule '%v' does not exist", rulespec)
|
require.Falsef(t, !exists && mustExists, "rule '%v' does not exist", rulespec)
|
||||||
require.Falsef(t, exists && !mustExists, "rule '%v' exist", rulespec)
|
require.Falsef(t, exists && !mustExists, "rule '%v' exist", rulespec)
|
||||||
@@ -283,7 +604,7 @@ func TestIptablesCreatePerformance(t *testing.T) {
|
|||||||
start := time.Now()
|
start := time.Now()
|
||||||
for i := 0; i < testMax; i++ {
|
for i := 0; i < testMax; i++ {
|
||||||
port := &fw.Port{Values: []uint16{uint16(1000 + i)}}
|
port := &fw.Port{Values: []uint16{uint16(1000 + i)}}
|
||||||
_, err = manager.AddPeerFiltering(nil, ip.AsSlice(), "tcp", nil, port, fw.ActionAccept, "")
|
_, err = manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, "tcp", nil, port, fw.ActionAccept)
|
||||||
|
|
||||||
require.NoError(t, err, "failed to add rule")
|
require.NoError(t, err, "failed to add rule")
|
||||||
}
|
}
|
||||||
@@ -291,40 +612,3 @@ func TestIptablesCreatePerformance(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestIptablesACLIPSetFallback verifies that when the kernel lacks ipset support,
|
|
||||||
// the ACL manager falls back to per-IP iptables rules (-s <ip>) instead of
|
|
||||||
// silently leaving the chain empty. See discussion #6125.
|
|
||||||
func TestIptablesACLIPSetFallback(t *testing.T) {
|
|
||||||
ipv4Client, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Use Create()/Init() so the router-owned chains (chainRTFWDIN/OUT) are
|
|
||||||
// created before the ACL manager's createDefaultChains() references them.
|
|
||||||
manager, err := Create(ifaceMock, iface.DefaultMTU)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, manager.Init(nil))
|
|
||||||
|
|
||||||
aclMgr := manager.aclMgr
|
|
||||||
// Simulate a kernel without the ipset hash module.
|
|
||||||
aclMgr.ipsetSupported = false
|
|
||||||
|
|
||||||
defer func() {
|
|
||||||
require.NoError(t, manager.Close(nil))
|
|
||||||
}()
|
|
||||||
|
|
||||||
ip := netip.MustParseAddr("10.20.0.42")
|
|
||||||
port := &fw.Port{Values: []uint16{22}}
|
|
||||||
|
|
||||||
rules, err := aclMgr.AddPeerFiltering(nil, ip.AsSlice(), "tcp", nil, port, fw.ActionAccept, "nb0000001")
|
|
||||||
require.NoError(t, err, "AddPeerFiltering should succeed via fallback")
|
|
||||||
require.NotEmpty(t, rules)
|
|
||||||
|
|
||||||
rule := rules[0].(*Rule)
|
|
||||||
require.Empty(t, rule.ipsetName, "fallback rule must not reference an ipset")
|
|
||||||
require.Contains(t, strings.Join(rule.specs, " "), "-s 10.20.0.42", "fallback rule must match by source IP")
|
|
||||||
require.NotContains(t, strings.Join(rule.specs, " "), "--match-set", "fallback rule must not use ipset matching")
|
|
||||||
|
|
||||||
// The rule must actually be present in the ACL chain (not silently dropped).
|
|
||||||
checkRuleSpecs(t, ipv4Client, rule.chain, true, rule.specs...)
|
|
||||||
}
|
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -31,7 +31,7 @@ func TestIptablesManager_RestoreOrCreateContainers(t *testing.T) {
|
|||||||
iptablesClient, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
iptablesClient, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
||||||
require.NoError(t, err, "failed to init iptables client")
|
require.NoError(t, err, "failed to init iptables client")
|
||||||
|
|
||||||
manager, err := newRouter(iptablesClient, ifaceMock, iface.DefaultMTU)
|
manager, err := newFamily(iptablesClient, ifaceMock, iface.DefaultMTU)
|
||||||
require.NoError(t, err, "should return a valid iptables manager")
|
require.NoError(t, err, "should return a valid iptables manager")
|
||||||
require.NoError(t, manager.init(nil))
|
require.NoError(t, manager.init(nil))
|
||||||
|
|
||||||
@@ -52,12 +52,12 @@ func TestIptablesManager_RestoreOrCreateContainers(t *testing.T) {
|
|||||||
// 11. MSS clamping rule for outbound traffic
|
// 11. MSS clamping rule for outbound traffic
|
||||||
require.Len(t, manager.rules, 11, "should have created rules map")
|
require.Len(t, manager.rules, 11, "should have created rules map")
|
||||||
|
|
||||||
exists, err := manager.iptablesClient.Exists(tableNat, chainPOSTROUTING, "-j", chainRTNAT)
|
exists, err := manager.iptablesClient.Exists(tableNat, chainPostrouting, "-j", chainRTNAT)
|
||||||
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableNat, chainPOSTROUTING)
|
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableNat, chainPostrouting)
|
||||||
require.True(t, exists, "postrouting jump rule should exist")
|
require.True(t, exists, "postrouting jump rule should exist")
|
||||||
|
|
||||||
exists, err = manager.iptablesClient.Exists(tableMangle, chainPREROUTING, "-j", chainRTPRE)
|
exists, err = manager.iptablesClient.Exists(tableMangle, chainPrerouting, "-j", chainRTPre)
|
||||||
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableMangle, chainPREROUTING)
|
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableMangle, chainPrerouting)
|
||||||
require.True(t, exists, "prerouting jump rule should exist")
|
require.True(t, exists, "prerouting jump rule should exist")
|
||||||
|
|
||||||
pair := firewall.RouterPair{
|
pair := firewall.RouterPair{
|
||||||
@@ -84,7 +84,7 @@ func TestIptablesManager_AddNatRule(t *testing.T) {
|
|||||||
iptablesClient, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
iptablesClient, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
||||||
require.NoError(t, err, "failed to init iptables client")
|
require.NoError(t, err, "failed to init iptables client")
|
||||||
|
|
||||||
manager, err := newRouter(iptablesClient, ifaceMock, iface.DefaultMTU)
|
manager, err := newFamily(iptablesClient, ifaceMock, iface.DefaultMTU)
|
||||||
require.NoError(t, err, "shouldn't return error")
|
require.NoError(t, err, "shouldn't return error")
|
||||||
require.NoError(t, manager.init(nil))
|
require.NoError(t, manager.init(nil))
|
||||||
|
|
||||||
@@ -95,7 +95,7 @@ func TestIptablesManager_AddNatRule(t *testing.T) {
|
|||||||
err = manager.AddNatRule(testCase.InputPair)
|
err = manager.AddNatRule(testCase.InputPair)
|
||||||
require.NoError(t, err, "marking rule should be inserted")
|
require.NoError(t, err, "marking rule should be inserted")
|
||||||
|
|
||||||
natRuleKey := firewall.GenKey(firewall.NatFormat, testCase.InputPair)
|
natRuleKey := testCase.InputPair.GenKey(firewall.NatFormat)
|
||||||
markingRule := []string{
|
markingRule := []string{
|
||||||
"-i", ifaceMock.Name(),
|
"-i", ifaceMock.Name(),
|
||||||
"-m", "conntrack",
|
"-m", "conntrack",
|
||||||
@@ -106,8 +106,8 @@ func TestIptablesManager_AddNatRule(t *testing.T) {
|
|||||||
fmt.Sprintf("%#x", nbnet.PreroutingFwmarkMasquerade),
|
fmt.Sprintf("%#x", nbnet.PreroutingFwmarkMasquerade),
|
||||||
}
|
}
|
||||||
|
|
||||||
exists, err := iptablesClient.Exists(tableMangle, chainRTPRE, markingRule...)
|
exists, err := iptablesClient.Exists(tableMangle, chainRTPre, markingRule...)
|
||||||
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableMangle, chainRTPRE)
|
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableMangle, chainRTPre)
|
||||||
if testCase.InputPair.Masquerade {
|
if testCase.InputPair.Masquerade {
|
||||||
require.True(t, exists, "marking rule should be created")
|
require.True(t, exists, "marking rule should be created")
|
||||||
foundRule, found := manager.rules[natRuleKey]
|
foundRule, found := manager.rules[natRuleKey]
|
||||||
@@ -121,7 +121,7 @@ func TestIptablesManager_AddNatRule(t *testing.T) {
|
|||||||
|
|
||||||
// Check inverse rule
|
// Check inverse rule
|
||||||
inversePair := firewall.GetInversePair(testCase.InputPair)
|
inversePair := firewall.GetInversePair(testCase.InputPair)
|
||||||
inverseRuleKey := firewall.GenKey(firewall.NatFormat, inversePair)
|
inverseRuleKey := inversePair.GenKey(firewall.NatFormat)
|
||||||
inverseMarkingRule := []string{
|
inverseMarkingRule := []string{
|
||||||
"!", "-i", ifaceMock.Name(),
|
"!", "-i", ifaceMock.Name(),
|
||||||
"-m", "conntrack",
|
"-m", "conntrack",
|
||||||
@@ -132,8 +132,8 @@ func TestIptablesManager_AddNatRule(t *testing.T) {
|
|||||||
fmt.Sprintf("%#x", nbnet.PreroutingFwmarkMasqueradeReturn),
|
fmt.Sprintf("%#x", nbnet.PreroutingFwmarkMasqueradeReturn),
|
||||||
}
|
}
|
||||||
|
|
||||||
exists, err = iptablesClient.Exists(tableMangle, chainRTPRE, inverseMarkingRule...)
|
exists, err = iptablesClient.Exists(tableMangle, chainRTPre, inverseMarkingRule...)
|
||||||
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableMangle, chainRTPRE)
|
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableMangle, chainRTPre)
|
||||||
if testCase.InputPair.Masquerade {
|
if testCase.InputPair.Masquerade {
|
||||||
require.True(t, exists, "inverse marking rule should be created")
|
require.True(t, exists, "inverse marking rule should be created")
|
||||||
foundRule, found := manager.rules[inverseRuleKey]
|
foundRule, found := manager.rules[inverseRuleKey]
|
||||||
@@ -157,7 +157,7 @@ func TestIptablesManager_RemoveNatRule(t *testing.T) {
|
|||||||
t.Run(testCase.Name, func(t *testing.T) {
|
t.Run(testCase.Name, func(t *testing.T) {
|
||||||
iptablesClient, _ := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
iptablesClient, _ := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
||||||
|
|
||||||
manager, err := newRouter(iptablesClient, ifaceMock, iface.DefaultMTU)
|
manager, err := newFamily(iptablesClient, ifaceMock, iface.DefaultMTU)
|
||||||
require.NoError(t, err, "shouldn't return error")
|
require.NoError(t, err, "shouldn't return error")
|
||||||
require.NoError(t, manager.init(nil))
|
require.NoError(t, manager.init(nil))
|
||||||
defer func() {
|
defer func() {
|
||||||
@@ -170,7 +170,7 @@ func TestIptablesManager_RemoveNatRule(t *testing.T) {
|
|||||||
err = manager.RemoveNatRule(testCase.InputPair)
|
err = manager.RemoveNatRule(testCase.InputPair)
|
||||||
require.NoError(t, err, "shouldn't return error")
|
require.NoError(t, err, "shouldn't return error")
|
||||||
|
|
||||||
natRuleKey := firewall.GenKey(firewall.NatFormat, testCase.InputPair)
|
natRuleKey := testCase.InputPair.GenKey(firewall.NatFormat)
|
||||||
markingRule := []string{
|
markingRule := []string{
|
||||||
"-i", ifaceMock.Name(),
|
"-i", ifaceMock.Name(),
|
||||||
"-m", "conntrack",
|
"-m", "conntrack",
|
||||||
@@ -181,8 +181,8 @@ func TestIptablesManager_RemoveNatRule(t *testing.T) {
|
|||||||
fmt.Sprintf("%#x", nbnet.PreroutingFwmarkMasquerade),
|
fmt.Sprintf("%#x", nbnet.PreroutingFwmarkMasquerade),
|
||||||
}
|
}
|
||||||
|
|
||||||
exists, err := iptablesClient.Exists(tableMangle, chainRTPRE, markingRule...)
|
exists, err := iptablesClient.Exists(tableMangle, chainRTPre, markingRule...)
|
||||||
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableMangle, chainRTPRE)
|
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableMangle, chainRTPre)
|
||||||
require.False(t, exists, "marking rule should not exist")
|
require.False(t, exists, "marking rule should not exist")
|
||||||
|
|
||||||
_, found := manager.rules[natRuleKey]
|
_, found := manager.rules[natRuleKey]
|
||||||
@@ -190,7 +190,7 @@ func TestIptablesManager_RemoveNatRule(t *testing.T) {
|
|||||||
|
|
||||||
// Check inverse rule removal
|
// Check inverse rule removal
|
||||||
inversePair := firewall.GetInversePair(testCase.InputPair)
|
inversePair := firewall.GetInversePair(testCase.InputPair)
|
||||||
inverseRuleKey := firewall.GenKey(firewall.NatFormat, inversePair)
|
inverseRuleKey := inversePair.GenKey(firewall.NatFormat)
|
||||||
inverseMarkingRule := []string{
|
inverseMarkingRule := []string{
|
||||||
"!", "-i", ifaceMock.Name(),
|
"!", "-i", ifaceMock.Name(),
|
||||||
"-m", "conntrack",
|
"-m", "conntrack",
|
||||||
@@ -201,8 +201,8 @@ func TestIptablesManager_RemoveNatRule(t *testing.T) {
|
|||||||
fmt.Sprintf("%#x", nbnet.PreroutingFwmarkMasqueradeReturn),
|
fmt.Sprintf("%#x", nbnet.PreroutingFwmarkMasqueradeReturn),
|
||||||
}
|
}
|
||||||
|
|
||||||
exists, err = iptablesClient.Exists(tableMangle, chainRTPRE, inverseMarkingRule...)
|
exists, err = iptablesClient.Exists(tableMangle, chainRTPre, inverseMarkingRule...)
|
||||||
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableMangle, chainRTPRE)
|
require.NoError(t, err, "should be able to query the iptables %s table and %s chain", tableMangle, chainRTPre)
|
||||||
require.False(t, exists, "inverse marking rule should not exist")
|
require.False(t, exists, "inverse marking rule should not exist")
|
||||||
|
|
||||||
_, found = manager.rules[inverseRuleKey]
|
_, found = manager.rules[inverseRuleKey]
|
||||||
@@ -219,13 +219,13 @@ func TestRouter_AddRouteFiltering(t *testing.T) {
|
|||||||
iptablesClient, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
iptablesClient, err := iptables.NewWithProtocol(iptables.ProtocolIPv4)
|
||||||
require.NoError(t, err, "Failed to create iptables client")
|
require.NoError(t, err, "Failed to create iptables client")
|
||||||
|
|
||||||
r, err := newRouter(iptablesClient, ifaceMock, iface.DefaultMTU)
|
r, err := newFamily(iptablesClient, ifaceMock, iface.DefaultMTU)
|
||||||
require.NoError(t, err, "Failed to create router manager")
|
require.NoError(t, err, "Failed to create family manager")
|
||||||
require.NoError(t, r.init(nil))
|
require.NoError(t, r.init(nil))
|
||||||
|
|
||||||
defer func() {
|
defer func() {
|
||||||
err := r.Reset()
|
err := r.Reset()
|
||||||
require.NoError(t, err, "Failed to reset router")
|
require.NoError(t, err, "Failed to reset family")
|
||||||
}()
|
}()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
@@ -334,62 +334,30 @@ func TestRouter_AddRouteFiltering(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
ruleKey, err := r.AddRouteFiltering(nil, tt.sources, firewall.Network{Prefix: tt.destination}, tt.proto, tt.sPort, tt.dPort, tt.action)
|
ruleKey, err := r.AddFilterRule(nil, tt.sources, firewall.Network{Prefix: tt.destination}, tt.proto, tt.sPort, tt.dPort, tt.action)
|
||||||
require.NoError(t, err, "AddRouteFiltering failed")
|
require.NoError(t, err, "AddFilterRule failed")
|
||||||
|
|
||||||
// Check if the rule is in the internal map
|
stored, ok := r.filters[ruleKey.ID()]
|
||||||
rule, ok := r.rules[ruleKey.ID()]
|
require.True(t, ok, "rule not stored in filters")
|
||||||
assert.True(t, ok, "Rule not found in internal map")
|
t.Logf("Internal rule: %v", stored.specs)
|
||||||
|
|
||||||
// Log the internal rule
|
exists, err := iptablesClient.Exists(tableFilter, chainRTFwdIn, stored.specs...)
|
||||||
t.Logf("Internal rule: %v", rule)
|
|
||||||
|
|
||||||
// Check if the rule exists in iptables
|
|
||||||
exists, err := iptablesClient.Exists(tableFilter, chainRTFWDIN, rule...)
|
|
||||||
assert.NoError(t, err, "Failed to check rule existence")
|
assert.NoError(t, err, "Failed to check rule existence")
|
||||||
assert.True(t, exists, "Rule not found in iptables")
|
assert.True(t, exists, "Rule not found in iptables")
|
||||||
|
|
||||||
var source firewall.Network
|
|
||||||
if len(tt.sources) > 1 {
|
|
||||||
source.Set = firewall.NewPrefixSet(tt.sources)
|
|
||||||
} else if len(tt.sources) > 0 {
|
|
||||||
source.Prefix = tt.sources[0]
|
|
||||||
}
|
|
||||||
// Verify rule content
|
|
||||||
params := routeFilteringRuleParams{
|
|
||||||
Source: source,
|
|
||||||
Destination: firewall.Network{Prefix: tt.destination},
|
|
||||||
Proto: tt.proto,
|
|
||||||
SPort: tt.sPort,
|
|
||||||
DPort: tt.dPort,
|
|
||||||
Action: tt.action,
|
|
||||||
}
|
|
||||||
|
|
||||||
expectedRule, err := r.genRouteRuleSpec(params, nil)
|
|
||||||
require.NoError(t, err, "Failed to generate expected rule spec")
|
|
||||||
|
|
||||||
if tt.expectSet {
|
if tt.expectSet {
|
||||||
setName := firewall.NewPrefixSet(tt.sources).HashedName()
|
setName := firewall.NewPrefixSet(tt.sources).HashedName()
|
||||||
expectedRule, err = r.genRouteRuleSpec(params, nil)
|
|
||||||
require.NoError(t, err, "Failed to generate expected rule spec with set")
|
|
||||||
|
|
||||||
// Check if the set was created
|
|
||||||
_, exists := r.ipsetCounter.Get(setName)
|
_, exists := r.ipsetCounter.Get(setName)
|
||||||
assert.True(t, exists, "IPSet not created")
|
assert.True(t, exists, "IPSet not created")
|
||||||
|
assert.NotEmpty(t, findSets(stored.specs), "Rule should reference an ipset")
|
||||||
}
|
}
|
||||||
|
|
||||||
assert.Equal(t, expectedRule, rule, "Rule content mismatch")
|
require.NoError(t, r.DeleteFilterRule(ruleKey), "Failed to delete rule")
|
||||||
|
|
||||||
// Clean up
|
|
||||||
err = r.DeleteRouteRule(ruleKey)
|
|
||||||
require.NoError(t, err, "Failed to delete rule")
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFindSetNameInRule(t *testing.T) {
|
func TestFindSetNameInRule(t *testing.T) {
|
||||||
r := &router{}
|
|
||||||
|
|
||||||
testCases := []struct {
|
testCases := []struct {
|
||||||
name string
|
name string
|
||||||
rule []string
|
rule []string
|
||||||
@@ -430,7 +398,7 @@ func TestFindSetNameInRule(t *testing.T) {
|
|||||||
|
|
||||||
for _, tc := range testCases {
|
for _, tc := range testCases {
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
result := r.findSets(tc.rule)
|
result := findSets(tc.rule)
|
||||||
|
|
||||||
if len(result) != len(tc.expected) {
|
if len(result) != len(tc.expected) {
|
||||||
t.Errorf("Expected %d sets, got %d. Sets found: %v", len(tc.expected), len(result), result)
|
t.Errorf("Expected %d sets, got %d. Sets found: %v", len(tc.expected), len(result), result)
|
||||||
|
|||||||
@@ -0,0 +1,273 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package iptables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
nbnet "github.com/netbirdio/netbird/client/net"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (r *family) AddNatRule(pair firewall.RouterPair) error {
|
||||||
|
if r.legacyManagement {
|
||||||
|
log.Warnf("This peer is connected to a NetBird Management service with an older version. Allowing all traffic for %s", pair.Destination)
|
||||||
|
if err := r.addLegacyRouteRule(pair); err != nil {
|
||||||
|
return fmt.Errorf("add legacy routing rule: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if pair.Masquerade {
|
||||||
|
if err := r.addNatRule(pair); err != nil {
|
||||||
|
return fmt.Errorf("add nat rule: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.addNatRule(firewall.GetInversePair(pair)); err != nil {
|
||||||
|
return fmt.Errorf("add inverse nat rule: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
r.updateState()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RemoveNatRule removes an iptables rule pair from forwarding and nat chains
|
||||||
|
func (r *family) RemoveNatRule(pair firewall.RouterPair) error {
|
||||||
|
if pair.Masquerade {
|
||||||
|
if err := r.removeNatRule(pair); err != nil {
|
||||||
|
return fmt.Errorf("remove nat rule: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.removeNatRule(firewall.GetInversePair(pair)); err != nil {
|
||||||
|
return fmt.Errorf("remove inverse nat rule: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.removeLegacyRouteRule(pair); err != nil {
|
||||||
|
return fmt.Errorf("remove legacy routing rule: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.updateState()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// addLegacyRouteRule adds a legacy routing rule for mgmt servers pre route acls
|
||||||
|
func (r *family) addLegacyRouteRule(pair firewall.RouterPair) error {
|
||||||
|
ruleID := pair.GenKey(firewall.ForwardingFormat)
|
||||||
|
|
||||||
|
if err := r.removeLegacyRouteRule(pair); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
rule := []string{"-s", pair.Source.String(), "-d", pair.Destination.String(), "-j", "ACCEPT"}
|
||||||
|
if err := r.iptablesClient.Append(tableFilter, chainRTFwdIn, rule...); err != nil {
|
||||||
|
return fmt.Errorf("add legacy forwarding rule %s -> %s: %w", pair.Source, pair.Destination, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.rules[ruleID] = rule
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) removeLegacyRouteRule(pair firewall.RouterPair) error {
|
||||||
|
ruleID := pair.GenKey(firewall.ForwardingFormat)
|
||||||
|
|
||||||
|
if rule, exists := r.rules[ruleID]; exists {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableFilter, chainRTFwdIn, rule...); err != nil {
|
||||||
|
return fmt.Errorf("remove legacy forwarding rule %s -> %s: %w", pair.Source, pair.Destination, err)
|
||||||
|
}
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
|
||||||
|
if err := r.decrementSetCounter(rule); err != nil {
|
||||||
|
return fmt.Errorf("decrement ipset counter: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetLegacyManagement returns the current legacy management mode
|
||||||
|
func (r *family) GetLegacyManagement() bool {
|
||||||
|
return r.legacyManagement
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetLegacyManagement sets the route manager to use legacy management mode
|
||||||
|
func (r *family) SetLegacyManagement(isLegacy bool) {
|
||||||
|
r.legacyManagement = isLegacy
|
||||||
|
}
|
||||||
|
|
||||||
|
// RemoveAllLegacyRouteRules removes all legacy routing rules for mgmt servers pre route acls
|
||||||
|
func (r *family) RemoveAllLegacyRouteRules() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
for k, rule := range r.rules {
|
||||||
|
if !strings.HasPrefix(string(k), firewall.ForwardingFormatPrefix) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableFilter, chainRTFwdIn, rule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove legacy forwarding rule: %w", err))
|
||||||
|
} else {
|
||||||
|
delete(r.rules, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
r.updateState()
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addPostroutingRules() error {
|
||||||
|
// First rule for outbound masquerade
|
||||||
|
rule1 := []string{
|
||||||
|
"-m", "mark", "--mark", fmt.Sprintf("%#x", nbnet.PreroutingFwmarkMasquerade),
|
||||||
|
"!", "-o", "lo",
|
||||||
|
"-j", "MASQUERADE",
|
||||||
|
}
|
||||||
|
if err := r.iptablesClient.Append(tableNat, chainRTNAT, rule1...); err != nil {
|
||||||
|
return fmt.Errorf("add outbound masquerade rule: %w", err)
|
||||||
|
}
|
||||||
|
r.rules["static-nat-outbound"] = rule1
|
||||||
|
|
||||||
|
// Second rule for return traffic masquerade
|
||||||
|
rule2 := []string{
|
||||||
|
"-m", "mark", "--mark", fmt.Sprintf("%#x", nbnet.PreroutingFwmarkMasqueradeReturn),
|
||||||
|
"-o", r.wgIface.Name(),
|
||||||
|
"-j", "MASQUERADE",
|
||||||
|
}
|
||||||
|
if err := r.iptablesClient.Append(tableNat, chainRTNAT, rule2...); err != nil {
|
||||||
|
return fmt.Errorf("add return masquerade rule: %w", err)
|
||||||
|
}
|
||||||
|
r.rules["static-nat-return"] = rule2
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// addMSSClampingRules adds MSS clamping rules to prevent fragmentation for forwarded traffic.
|
||||||
|
func (r *family) addMSSClampingRules() error {
|
||||||
|
overhead := uint16(ipv4TCPHeaderSize)
|
||||||
|
if r.v6 {
|
||||||
|
overhead = ipv6TCPHeaderSize
|
||||||
|
}
|
||||||
|
mss := r.mtu - overhead
|
||||||
|
|
||||||
|
// Add jump rule from FORWARD chain in mangle table to our custom chain
|
||||||
|
jumpRule := jumpRuleSpec(chainRTMSSClamp)
|
||||||
|
if err := r.iptablesClient.Insert(tableMangle, chainForward, 1, jumpRule...); err != nil {
|
||||||
|
return fmt.Errorf("add jump to MSS clamp chain: %w", err)
|
||||||
|
}
|
||||||
|
r.rules[jumpMSSClamp] = jumpRule
|
||||||
|
|
||||||
|
ruleOut := []string{
|
||||||
|
"-o", r.wgIface.Name(),
|
||||||
|
"-p", "tcp",
|
||||||
|
"--tcp-flags", "SYN,RST", "SYN",
|
||||||
|
"-j", "TCPMSS",
|
||||||
|
"--set-mss", fmt.Sprintf("%d", mss),
|
||||||
|
}
|
||||||
|
if err := r.iptablesClient.Append(tableMangle, chainRTMSSClamp, ruleOut...); err != nil {
|
||||||
|
return fmt.Errorf("add outbound MSS clamp rule: %w", err)
|
||||||
|
}
|
||||||
|
r.rules["mss-clamp-out"] = ruleOut
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) insertEstablishedRule(chain string) error {
|
||||||
|
establishedRule := getConntrackEstablished()
|
||||||
|
|
||||||
|
err := r.iptablesClient.Insert(tableFilter, chain, 1, establishedRule...)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("insert established rule: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ruleID := firewall.RuleID("established-" + chain)
|
||||||
|
r.rules[ruleID] = establishedRule
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addNatRule(pair firewall.RouterPair) (err error) {
|
||||||
|
ruleID := pair.GenKey(firewall.NatFormat)
|
||||||
|
|
||||||
|
if rule, exists := r.rules[ruleID]; exists {
|
||||||
|
if derr := r.iptablesClient.DeleteIfExists(tableMangle, chainRTPre, rule...); derr != nil {
|
||||||
|
return fmt.Errorf("remove existing marking rule for %s: %w", pair.Destination, derr)
|
||||||
|
}
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
|
||||||
|
// Drop the replaced spec's set references only once the new spec has
|
||||||
|
// taken its own, so a set both specs share is not destroyed and
|
||||||
|
// recreated, which would lose the prefixes UpdateSet put in it.
|
||||||
|
defer func() {
|
||||||
|
if derr := r.decrementSetCounter(rule); derr != nil && err == nil {
|
||||||
|
err = fmt.Errorf("decrement ipset counter: %w", derr)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
markValue := nbnet.PreroutingFwmarkMasquerade
|
||||||
|
if pair.Inverse {
|
||||||
|
markValue = nbnet.PreroutingFwmarkMasqueradeReturn
|
||||||
|
}
|
||||||
|
|
||||||
|
rule := []string{"-i", r.wgIface.Name()}
|
||||||
|
if pair.Inverse {
|
||||||
|
rule = []string{"!", "-i", r.wgIface.Name()}
|
||||||
|
}
|
||||||
|
|
||||||
|
rule = append(rule,
|
||||||
|
"-m", "conntrack",
|
||||||
|
"--ctstate", "NEW",
|
||||||
|
)
|
||||||
|
sourceExp, err := r.applyNetwork("-s", pair.Source, nil)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("apply network -s: %w", err)
|
||||||
|
}
|
||||||
|
destExp, err := r.applyNetwork("-d", pair.Destination, nil)
|
||||||
|
if err != nil {
|
||||||
|
r.dropSourceMatch(sourceExp)
|
||||||
|
return fmt.Errorf("apply network -d: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rule = append(rule, sourceExp...)
|
||||||
|
rule = append(rule, destExp...)
|
||||||
|
rule = append(rule,
|
||||||
|
"-j", "MARK", "--set-mark", fmt.Sprintf("%#x", markValue),
|
||||||
|
)
|
||||||
|
|
||||||
|
// Ensure nat rules come first, so the mark can be overwritten.
|
||||||
|
// Currently overwritten by the dst-type LOCAL rules for redirected traffic.
|
||||||
|
if err := r.iptablesClient.Insert(tableMangle, chainRTPre, 1, rule...); err != nil {
|
||||||
|
r.dropSourceMatch(rule)
|
||||||
|
return fmt.Errorf("add marking rule for %s: %w", pair.Destination, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.rules[ruleID] = rule
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) removeNatRule(pair firewall.RouterPair) error {
|
||||||
|
ruleID := pair.GenKey(firewall.NatFormat)
|
||||||
|
|
||||||
|
if rule, exists := r.rules[ruleID]; exists {
|
||||||
|
if err := r.iptablesClient.DeleteIfExists(tableMangle, chainRTPre, rule...); err != nil {
|
||||||
|
return fmt.Errorf("remove marking rule for %s: %w", pair.Destination, err)
|
||||||
|
}
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
|
||||||
|
if err := r.decrementSetCounter(rule); err != nil {
|
||||||
|
return fmt.Errorf("decrement ipset counter: %w", err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
log.Debugf("marking rule %s not found", ruleID)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -1,18 +1,37 @@
|
|||||||
package iptables
|
package iptables
|
||||||
|
|
||||||
// Rule to handle management of rules
|
import "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
type Rule struct {
|
|
||||||
ruleID string
|
|
||||||
ipsetName string
|
|
||||||
|
|
||||||
|
// Rule to handle management of rules. Source set membership (when the
|
||||||
|
// rule was built against a shared hash:net ipset) is encoded in specs;
|
||||||
|
// DeleteFilterRule recovers it via findSets so the refcounter can drop
|
||||||
|
// the right reference.
|
||||||
|
type Rule struct {
|
||||||
|
id manager.RuleID
|
||||||
specs []string
|
specs []string
|
||||||
mangleSpecs []string
|
mangleSpecs []string
|
||||||
ip string
|
// extraRules holds the rules beyond the first when the ipset
|
||||||
chain string
|
// fallback expands a multi-source rule into one rule per prefix.
|
||||||
v6 bool
|
extraRules []filterSpecs
|
||||||
|
chain string
|
||||||
|
v6 bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetRuleID returns the rule id
|
// filterSpecs is one installed iptables rule: its filter-table spec and
|
||||||
func (r *Rule) ID() string {
|
// the paired mangle redirect-mark spec (nil for route rules or when the
|
||||||
return r.ruleID
|
// mangle rule could not be added).
|
||||||
|
type filterSpecs struct {
|
||||||
|
specs []string
|
||||||
|
mangleSpecs []string
|
||||||
|
}
|
||||||
|
|
||||||
|
// allSpecs returns the spec pairs of every iptables rule backing this
|
||||||
|
// Rule, the primary one first.
|
||||||
|
func (r *Rule) allSpecs() []filterSpecs {
|
||||||
|
return append([]filterSpecs{{specs: r.specs, mangleSpecs: r.mangleSpecs}}, r.extraRules...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ID returns the rule id
|
||||||
|
func (r *Rule) ID() manager.RuleID {
|
||||||
|
return r.id
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,127 +0,0 @@
|
|||||||
package iptables
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/json"
|
|
||||||
"maps"
|
|
||||||
)
|
|
||||||
|
|
||||||
type ipList struct {
|
|
||||||
ips map[string]struct{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newIpList(ip string) *ipList {
|
|
||||||
ips := make(map[string]struct{})
|
|
||||||
ips[ip] = struct{}{}
|
|
||||||
|
|
||||||
return &ipList{
|
|
||||||
ips: ips,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipList) addIP(ip string) {
|
|
||||||
s.ips[ip] = struct{}{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// clone returns a deep copy of the ipList with its own ips map.
|
|
||||||
func (s *ipList) clone() *ipList {
|
|
||||||
if s == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return &ipList{ips: maps.Clone(s.ips)}
|
|
||||||
}
|
|
||||||
|
|
||||||
// MarshalJSON implements json.Marshaler
|
|
||||||
func (s *ipList) MarshalJSON() ([]byte, error) {
|
|
||||||
return json.Marshal(struct {
|
|
||||||
IPs map[string]struct{} `json:"ips"`
|
|
||||||
}{
|
|
||||||
IPs: s.ips,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// UnmarshalJSON implements json.Unmarshaler
|
|
||||||
func (s *ipList) UnmarshalJSON(data []byte) error {
|
|
||||||
temp := struct {
|
|
||||||
IPs map[string]struct{} `json:"ips"`
|
|
||||||
}{}
|
|
||||||
if err := json.Unmarshal(data, &temp); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
s.ips = temp.IPs
|
|
||||||
|
|
||||||
if temp.IPs == nil {
|
|
||||||
temp.IPs = make(map[string]struct{})
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type ipsetStore struct {
|
|
||||||
ipsets map[string]*ipList
|
|
||||||
}
|
|
||||||
|
|
||||||
func newIpsetStore() *ipsetStore {
|
|
||||||
return &ipsetStore{
|
|
||||||
ipsets: make(map[string]*ipList),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// clone returns a deep copy of the ipsetStore with its own ipsets map and
|
|
||||||
// independent ipList entries.
|
|
||||||
func (s *ipsetStore) clone() *ipsetStore {
|
|
||||||
if s == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
cloned := &ipsetStore{ipsets: make(map[string]*ipList, len(s.ipsets))}
|
|
||||||
for name, list := range s.ipsets {
|
|
||||||
cloned.ipsets[name] = list.clone()
|
|
||||||
}
|
|
||||||
return cloned
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) ipset(ipsetName string) (*ipList, bool) {
|
|
||||||
r, ok := s.ipsets[ipsetName]
|
|
||||||
return r, ok
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) addIpList(ipsetName string, list *ipList) {
|
|
||||||
s.ipsets[ipsetName] = list
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) deleteIpset(ipsetName string) {
|
|
||||||
delete(s.ipsets, ipsetName)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) ipsetNames() []string {
|
|
||||||
names := make([]string, 0, len(s.ipsets))
|
|
||||||
for name := range s.ipsets {
|
|
||||||
names = append(names, name)
|
|
||||||
}
|
|
||||||
return names
|
|
||||||
}
|
|
||||||
|
|
||||||
// MarshalJSON implements json.Marshaler
|
|
||||||
func (s *ipsetStore) MarshalJSON() ([]byte, error) {
|
|
||||||
return json.Marshal(struct {
|
|
||||||
IPSets map[string]*ipList `json:"ipsets"`
|
|
||||||
}{
|
|
||||||
IPSets: s.ipsets,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// UnmarshalJSON implements json.Unmarshaler
|
|
||||||
func (s *ipsetStore) UnmarshalJSON(data []byte) error {
|
|
||||||
temp := struct {
|
|
||||||
IPSets map[string]*ipList `json:"ipsets"`
|
|
||||||
}{}
|
|
||||||
if err := json.Unmarshal(data, &temp); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
s.ipsets = temp.IPSets
|
|
||||||
|
|
||||||
if temp.IPSets == nil {
|
|
||||||
temp.IPSets = make(map[string]*ipList)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -29,17 +29,13 @@ type ShutdownState struct {
|
|||||||
|
|
||||||
InterfaceState *InterfaceState `json:"interface_state,omitempty"`
|
InterfaceState *InterfaceState `json:"interface_state,omitempty"`
|
||||||
|
|
||||||
RouteRules routeRules `json:"route_rules,omitempty"`
|
RouteRules routeRules `json:"route_rules,omitempty"`
|
||||||
RouteIPsetCounter *ipsetCounter `json:"route_ipset_counter,omitempty"`
|
|
||||||
|
|
||||||
ACLEntries aclEntries `json:"acl_entries,omitempty"`
|
|
||||||
ACLIPsetStore *ipsetStore `json:"acl_ipset_store,omitempty"`
|
|
||||||
|
|
||||||
// IPv6 counterparts
|
|
||||||
RouteRules6 routeRules `json:"route_rules_v6,omitempty"`
|
RouteRules6 routeRules `json:"route_rules_v6,omitempty"`
|
||||||
|
RouteIPsetCounter *ipsetCounter `json:"route_ipset_counter,omitempty"`
|
||||||
RouteIPsetCounter6 *ipsetCounter `json:"route_ipset_counter_v6,omitempty"`
|
RouteIPsetCounter6 *ipsetCounter `json:"route_ipset_counter_v6,omitempty"`
|
||||||
ACLEntries6 aclEntries `json:"acl_entries_v6,omitempty"`
|
|
||||||
ACLIPsetStore6 *ipsetStore `json:"acl_ipset_store_v6,omitempty"`
|
ACLEntries aclEntries `json:"acl_entries,omitempty"`
|
||||||
|
ACLEntries6 aclEntries `json:"acl_entries_v6,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *ShutdownState) Name() string {
|
func (s *ShutdownState) Name() string {
|
||||||
@@ -57,17 +53,14 @@ func (s *ShutdownState) Cleanup() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if s.RouteRules != nil {
|
if s.RouteRules != nil {
|
||||||
ipt.router.rules = s.RouteRules
|
ipt.family4.rules = s.RouteRules
|
||||||
}
|
}
|
||||||
if s.RouteIPsetCounter != nil {
|
if s.RouteIPsetCounter != nil {
|
||||||
ipt.router.ipsetCounter.LoadData(s.RouteIPsetCounter)
|
ipt.family4.ipsetCounter.LoadData(s.RouteIPsetCounter)
|
||||||
}
|
}
|
||||||
|
|
||||||
if s.ACLEntries != nil {
|
if s.ACLEntries != nil {
|
||||||
ipt.aclMgr.entries = s.ACLEntries
|
ipt.family4.entries = s.ACLEntries
|
||||||
}
|
|
||||||
if s.ACLIPsetStore != nil {
|
|
||||||
ipt.aclMgr.ipsetStore = s.ACLIPsetStore
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Clean up v6 state even if the current run has no IPv6.
|
// Clean up v6 state even if the current run has no IPv6.
|
||||||
@@ -79,16 +72,13 @@ func (s *ShutdownState) Cleanup() error {
|
|||||||
}
|
}
|
||||||
if ipt.hasIPv6() {
|
if ipt.hasIPv6() {
|
||||||
if s.RouteRules6 != nil {
|
if s.RouteRules6 != nil {
|
||||||
ipt.router6.rules = s.RouteRules6
|
ipt.family6.rules = s.RouteRules6
|
||||||
}
|
}
|
||||||
if s.RouteIPsetCounter6 != nil {
|
if s.RouteIPsetCounter6 != nil {
|
||||||
ipt.router6.ipsetCounter.LoadData(s.RouteIPsetCounter6)
|
ipt.family6.ipsetCounter.LoadData(s.RouteIPsetCounter6)
|
||||||
}
|
}
|
||||||
if s.ACLEntries6 != nil {
|
if s.ACLEntries6 != nil {
|
||||||
ipt.aclMgr6.entries = s.ACLEntries6
|
ipt.family6.entries = s.ACLEntries6
|
||||||
}
|
|
||||||
if s.ACLIPsetStore6 != nil {
|
|
||||||
ipt.aclMgr6.ipsetStore = s.ACLIPsetStore6
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,27 @@
|
|||||||
|
//go:build privileged
|
||||||
|
|
||||||
|
package iptables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
)
|
||||||
|
|
||||||
|
func pfx(ip net.IP) []netip.Prefix {
|
||||||
|
if ip == nil {
|
||||||
|
return []netip.Prefix{netip.PrefixFrom(netip.IPv4Unspecified(), 0)}
|
||||||
|
}
|
||||||
|
if ip.IsUnspecified() {
|
||||||
|
if ip.To4() != nil {
|
||||||
|
return []netip.Prefix{netip.PrefixFrom(netip.IPv4Unspecified(), 0)}
|
||||||
|
}
|
||||||
|
return []netip.Prefix{netip.PrefixFrom(netip.IPv6Unspecified(), 0)}
|
||||||
|
}
|
||||||
|
a, ok := netip.AddrFromSlice(ip)
|
||||||
|
if !ok {
|
||||||
|
panic(fmt.Sprintf("invalid IP length: %d", len(ip)))
|
||||||
|
}
|
||||||
|
a = a.Unmap()
|
||||||
|
return []netip.Prefix{netip.PrefixFrom(a, a.BitLen())}
|
||||||
|
}
|
||||||
@@ -3,7 +3,6 @@ package manager
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"sort"
|
"sort"
|
||||||
|
|
||||||
@@ -16,6 +15,12 @@ import (
|
|||||||
// method but the IPv6 firewall components were not initialized.
|
// method but the IPv6 firewall components were not initialized.
|
||||||
var ErrIPv6NotInitialized = errors.New("IPv6 firewall not initialized")
|
var ErrIPv6NotInitialized = errors.New("IPv6 firewall not initialized")
|
||||||
|
|
||||||
|
// ErrNoSources is returned when AddFilterRule is called with an empty
|
||||||
|
// source list. "Match any source" must be expressed explicitly with a
|
||||||
|
// /0 prefix; an empty list is a caller error and is rejected rather
|
||||||
|
// than silently widening the rule to every source.
|
||||||
|
var ErrNoSources = errors.New("rule has no sources")
|
||||||
|
|
||||||
const (
|
const (
|
||||||
ForwardingFormatPrefix = "netbird-fwd-"
|
ForwardingFormatPrefix = "netbird-fwd-"
|
||||||
ForwardingFormat = "netbird-fwd-%s-%t"
|
ForwardingFormat = "netbird-fwd-%s-%t"
|
||||||
@@ -23,13 +28,18 @@ const (
|
|||||||
NatFormat = "netbird-nat-%s-%t"
|
NatFormat = "netbird-nat-%s-%t"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// RuleID identifies a firewall rule. It is a typed string so the
|
||||||
|
// compiler catches accidental mixing with arbitrary string keys. It is
|
||||||
|
// only an identifier and does not implement Rule.
|
||||||
|
type RuleID string
|
||||||
|
|
||||||
// Rule abstraction should be implemented by each firewall manager
|
// Rule abstraction should be implemented by each firewall manager
|
||||||
//
|
//
|
||||||
// Each firewall type for different OS can use different type
|
// Each firewall type for different OS can use different type
|
||||||
// of the properties to hold data of the created rule
|
// of the properties to hold data of the created rule
|
||||||
type Rule interface {
|
type Rule interface {
|
||||||
// ID returns the rule id
|
// ID returns the rule id
|
||||||
ID() string
|
ID() RuleID
|
||||||
}
|
}
|
||||||
|
|
||||||
// RuleDirection is the traffic direction which a rule is applied
|
// RuleDirection is the traffic direction which a rule is applied
|
||||||
@@ -91,6 +101,13 @@ func (d Network) IsPrefix() bool {
|
|||||||
return d.Prefix.IsValid()
|
return d.Prefix.IsValid()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// IsZero returns true if the network designates no destination, i.e. it
|
||||||
|
// is the zero value. A zero Network is the peer-rule sentinel; a non-zero
|
||||||
|
// one carries a prefix or set destination.
|
||||||
|
func (d Network) IsZero() bool {
|
||||||
|
return !d.IsPrefix() && !d.IsSet()
|
||||||
|
}
|
||||||
|
|
||||||
// Manager is the high level abstraction of a firewall manager
|
// Manager is the high level abstraction of a firewall manager
|
||||||
//
|
//
|
||||||
// It declares methods which handle actions required by the
|
// It declares methods which handle actions required by the
|
||||||
@@ -98,46 +115,42 @@ func (d Network) IsPrefix() bool {
|
|||||||
type Manager interface {
|
type Manager interface {
|
||||||
Init(stateManager *statemanager.Manager) error
|
Init(stateManager *statemanager.Manager) error
|
||||||
|
|
||||||
// AllowNetbird allows netbird interface traffic
|
// AddFilterRule adds a packet-filtering rule to the firewall.
|
||||||
AllowNetbird() error
|
|
||||||
|
|
||||||
// AddPeerFiltering adds a rule to the firewall
|
|
||||||
//
|
//
|
||||||
// If comment argument is empty firewall manager should set
|
// If destination is the zero Network, the rule applies to traffic
|
||||||
// rule ID as comment for the rule
|
// inbound to this node, i.e. peer ACL semantics, installed in
|
||||||
|
// the kernel's input chain. If destination is set (prefix or
|
||||||
|
// set), the rule applies to forwarded traffic with that
|
||||||
|
// destination, route ACL semantics, installed in the forward
|
||||||
|
// chain.
|
||||||
//
|
//
|
||||||
// Note: Callers should call Flush() after adding rules to ensure
|
// sources must be a single address family; the caller splits mixed
|
||||||
// they are applied to the kernel and rule handles are refreshed.
|
// families and calls once per family. "Match any source" must be
|
||||||
AddPeerFiltering(
|
// expressed with an explicit /0 prefix; an empty sources list is
|
||||||
|
// rejected with ErrNoSources so a zeroed list can never widen a
|
||||||
|
// rule to every source.
|
||||||
|
//
|
||||||
|
// Note: callers should call Flush() after adding rules.
|
||||||
|
AddFilterRule(
|
||||||
id []byte,
|
id []byte,
|
||||||
ip net.IP,
|
sources []netip.Prefix,
|
||||||
|
destination Network,
|
||||||
proto Protocol,
|
proto Protocol,
|
||||||
sPort *Port,
|
sPort *Port,
|
||||||
dPort *Port,
|
dPort *Port,
|
||||||
action Action,
|
action Action,
|
||||||
ipsetName string,
|
) (Rule, error)
|
||||||
) ([]Rule, error)
|
|
||||||
|
|
||||||
// DeletePeerRule from the firewall by rule definition
|
// DeleteFilterRule removes a filtering rule previously added via
|
||||||
DeletePeerRule(rule Rule) error
|
// AddFilterRule. The rule's own type identifies whether it lives
|
||||||
|
// in the peer (input) or route (forward) path.
|
||||||
|
DeleteFilterRule(rule Rule) error
|
||||||
|
|
||||||
// IsServerRouteSupported returns true if the firewall supports server side routing operations
|
// IsServerRouteSupported returns true if the firewall supports server side routing operations
|
||||||
IsServerRouteSupported() bool
|
IsServerRouteSupported() bool
|
||||||
|
|
||||||
IsStateful() bool
|
IsStateful() bool
|
||||||
|
|
||||||
AddRouteFiltering(
|
|
||||||
id []byte,
|
|
||||||
sources []netip.Prefix,
|
|
||||||
destination Network,
|
|
||||||
proto Protocol,
|
|
||||||
sPort, dPort *Port,
|
|
||||||
action Action,
|
|
||||||
) (Rule, error)
|
|
||||||
|
|
||||||
// DeleteRouteRule deletes a routing rule
|
|
||||||
DeleteRouteRule(rule Rule) error
|
|
||||||
|
|
||||||
// AddNatRule inserts a routing NAT rule
|
// AddNatRule inserts a routing NAT rule
|
||||||
AddNatRule(pair RouterPair) error
|
AddNatRule(pair RouterPair) error
|
||||||
|
|
||||||
@@ -179,14 +192,11 @@ 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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func GenKey(format string, pair RouterPair) string {
|
// GenKey builds the rule id for this pair from the given format.
|
||||||
return fmt.Sprintf(format, pair.ID, pair.Inverse)
|
func (p RouterPair) GenKey(format string) RuleID {
|
||||||
|
return RuleID(fmt.Sprintf(format, p.ID, p.Inverse))
|
||||||
}
|
}
|
||||||
|
|
||||||
// LegacyManager defines the interface for legacy management operations
|
// LegacyManager defines the interface for legacy management operations
|
||||||
@@ -242,6 +252,20 @@ func MergeIPRanges(prefixes []netip.Prefix) []netip.Prefix {
|
|||||||
return merged
|
return merged
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// UnmapPrefix normalizes a v4-mapped v6 prefix (::ffff:a.b.c.d) to its
|
||||||
|
// plain v4 form, shifting the prefix length out of the 96-bit mapped
|
||||||
|
// range. Other prefixes are returned unchanged. Keeping prefixes
|
||||||
|
// unmapped ensures v4 rules match consistently and the match builders
|
||||||
|
// read the correct address length.
|
||||||
|
func UnmapPrefix(p netip.Prefix) netip.Prefix {
|
||||||
|
addr := p.Addr()
|
||||||
|
if !addr.Is4In6() {
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
bits := max(p.Bits()-96, 0)
|
||||||
|
return netip.PrefixFrom(addr.Unmap(), bits)
|
||||||
|
}
|
||||||
|
|
||||||
// SortPrefixes sorts the given slice of netip.Prefix in place.
|
// SortPrefixes sorts the given slice of netip.Prefix in place.
|
||||||
// It sorts first by IP address, then by prefix length (most specific to least specific).
|
// It sorts first by IP address, then by prefix length (most specific to least specific).
|
||||||
func SortPrefixes(prefixes []netip.Prefix) {
|
func SortPrefixes(prefixes []netip.Prefix) {
|
||||||
|
|||||||
@@ -13,13 +13,13 @@ type ForwardRule struct {
|
|||||||
TranslatedPort Port
|
TranslatedPort Port
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r ForwardRule) ID() string {
|
func (r ForwardRule) ID() RuleID {
|
||||||
id := fmt.Sprintf("%s;%s;%s;%s",
|
id := fmt.Sprintf("%s;%s;%s;%s",
|
||||||
r.Protocol,
|
r.Protocol,
|
||||||
r.DestinationPort.String(),
|
r.DestinationPort.String(),
|
||||||
r.TranslatedAddress.String(),
|
r.TranslatedAddress.String(),
|
||||||
r.TranslatedPort.String())
|
r.TranslatedPort.String())
|
||||||
return id
|
return RuleID(id)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r ForwardRule) String() string {
|
func (r ForwardRule) String() string {
|
||||||
|
|||||||
@@ -40,7 +40,7 @@ func (h Set) Comment() string {
|
|||||||
|
|
||||||
// NewPrefixSet generates a unique name for an ipset based on the given prefixes.
|
// NewPrefixSet generates a unique name for an ipset based on the given prefixes.
|
||||||
func NewPrefixSet(prefixes []netip.Prefix) Set {
|
func NewPrefixSet(prefixes []netip.Prefix) Set {
|
||||||
// sort for consistent naming
|
prefixes = slices.Clone(prefixes)
|
||||||
SortPrefixes(prefixes)
|
SortPrefixes(prefixes)
|
||||||
|
|
||||||
hash := sha256.New()
|
hash := sha256.New()
|
||||||
|
|||||||
@@ -1,713 +0,0 @@
|
|||||||
package nftables
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"fmt"
|
|
||||||
"net"
|
|
||||||
"slices"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/google/nftables"
|
|
||||||
"github.com/google/nftables/binaryutil"
|
|
||||||
"github.com/google/nftables/expr"
|
|
||||||
log "github.com/sirupsen/logrus"
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
|
|
||||||
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
|
||||||
nbnet "github.com/netbirdio/netbird/client/net"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
|
|
||||||
// rules chains contains the effective ACL rules
|
|
||||||
chainNameInputRules = "netbird-acl-input-rules"
|
|
||||||
|
|
||||||
// filter chains contains the rules that jump to the rules chains
|
|
||||||
chainNameInputFilter = "netbird-acl-input-filter"
|
|
||||||
chainNameForwardFilter = "netbird-acl-forward-filter"
|
|
||||||
chainNameManglePrerouting = "netbird-mangle-prerouting"
|
|
||||||
chainNameManglePostrouting = "netbird-mangle-postrouting"
|
|
||||||
)
|
|
||||||
|
|
||||||
const flushError = "flush: %w"
|
|
||||||
|
|
||||||
type AclManager struct {
|
|
||||||
rConn *nftables.Conn
|
|
||||||
sConn *nftables.Conn
|
|
||||||
wgIface iFaceMapper
|
|
||||||
routingFwChainName string
|
|
||||||
af addrFamily
|
|
||||||
|
|
||||||
workTable *nftables.Table
|
|
||||||
chainInputRules *nftables.Chain
|
|
||||||
chainPrerouting *nftables.Chain
|
|
||||||
|
|
||||||
ipsetStore *ipsetStore
|
|
||||||
rules map[string]*Rule
|
|
||||||
}
|
|
||||||
|
|
||||||
func newAclManager(table *nftables.Table, wgIface iFaceMapper, routingFwChainName string) (*AclManager, error) {
|
|
||||||
// sConn is used for creating sets and adding/removing elements from them
|
|
||||||
// it's differ then rConn (which does create new conn for each flush operation)
|
|
||||||
// and is permanent. Using same connection for both type of operations
|
|
||||||
// overloads netlink with high amount of rules ( > 10000)
|
|
||||||
sConn, err := nftables.New(nftables.AsLasting())
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("create nf conn: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return &AclManager{
|
|
||||||
rConn: &nftables.Conn{},
|
|
||||||
sConn: sConn,
|
|
||||||
wgIface: wgIface,
|
|
||||||
workTable: table,
|
|
||||||
routingFwChainName: routingFwChainName,
|
|
||||||
af: familyForAddr(table.Family == nftables.TableFamilyIPv4),
|
|
||||||
|
|
||||||
ipsetStore: newIpsetStore(),
|
|
||||||
rules: make(map[string]*Rule),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) init(workTable *nftables.Table) error {
|
|
||||||
m.workTable = workTable
|
|
||||||
return m.createDefaultChains()
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddPeerFiltering rule to the firewall
|
|
||||||
//
|
|
||||||
// If comment argument is empty firewall manager should set
|
|
||||||
// rule ID as comment for the rule
|
|
||||||
func (m *AclManager) AddPeerFiltering(
|
|
||||||
id []byte,
|
|
||||||
ip net.IP,
|
|
||||||
proto firewall.Protocol,
|
|
||||||
sPort *firewall.Port,
|
|
||||||
dPort *firewall.Port,
|
|
||||||
action firewall.Action,
|
|
||||||
ipsetName string,
|
|
||||||
) ([]firewall.Rule, error) {
|
|
||||||
var ipset *nftables.Set
|
|
||||||
if ipsetName != "" {
|
|
||||||
var err error
|
|
||||||
ipset, err = m.addIpToSet(ipsetName, ip)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
newRules := make([]firewall.Rule, 0, 2)
|
|
||||||
ioRule, err := m.addIOFiltering(ip, proto, sPort, dPort, action, ipset)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
newRules = append(newRules, ioRule)
|
|
||||||
return newRules, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// DeletePeerRule from the firewall by rule definition
|
|
||||||
func (m *AclManager) DeletePeerRule(rule firewall.Rule) error {
|
|
||||||
r, ok := rule.(*Rule)
|
|
||||||
if !ok {
|
|
||||||
return fmt.Errorf("invalid rule type")
|
|
||||||
}
|
|
||||||
|
|
||||||
if r.nftSet == nil {
|
|
||||||
if err := m.rConn.DelRule(r.nftRule); err != nil {
|
|
||||||
log.Errorf("failed to delete rule: %v", err)
|
|
||||||
}
|
|
||||||
if r.mangleRule != nil {
|
|
||||||
if err := m.rConn.DelRule(r.mangleRule); err != nil {
|
|
||||||
log.Errorf("failed to delete mangle rule: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
delete(m.rules, r.ID())
|
|
||||||
return m.rConn.Flush()
|
|
||||||
}
|
|
||||||
|
|
||||||
ips, ok := m.ipsetStore.ips(r.nftSet.Name)
|
|
||||||
if !ok {
|
|
||||||
if err := m.rConn.DelRule(r.nftRule); err != nil {
|
|
||||||
log.Errorf("failed to delete rule: %v", err)
|
|
||||||
}
|
|
||||||
if r.mangleRule != nil {
|
|
||||||
if err := m.rConn.DelRule(r.mangleRule); err != nil {
|
|
||||||
log.Errorf("failed to delete mangle rule: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
delete(m.rules, r.ID())
|
|
||||||
return m.rConn.Flush()
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, ok := ips[r.ip.String()]; ok {
|
|
||||||
err := m.sConn.SetDeleteElements(r.nftSet, []nftables.SetElement{{Key: ipToBytes(r.ip, m.af)}})
|
|
||||||
if err != nil {
|
|
||||||
log.Errorf("delete elements for set %q: %v", r.nftSet.Name, err)
|
|
||||||
}
|
|
||||||
if err := m.sConn.Flush(); err != nil {
|
|
||||||
log.Debugf("flush error of set delete element, %s", r.nftSet.Name)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
m.ipsetStore.DeleteIpFromSet(r.nftSet.Name, r.ip)
|
|
||||||
}
|
|
||||||
|
|
||||||
// if after delete, set still contains other IPs,
|
|
||||||
// no need to delete firewall rule and we should exit here
|
|
||||||
if len(ips) > 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.rConn.DelRule(r.nftRule); err != nil {
|
|
||||||
log.Errorf("failed to delete rule: %v", err)
|
|
||||||
}
|
|
||||||
if r.mangleRule != nil {
|
|
||||||
if err := m.rConn.DelRule(r.mangleRule); err != nil {
|
|
||||||
log.Errorf("failed to delete mangle rule: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.rConn.Flush(); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
delete(m.rules, r.ID())
|
|
||||||
m.ipsetStore.DeleteReferenceFromIpSet(r.nftSet.Name)
|
|
||||||
|
|
||||||
if m.ipsetStore.HasReferenceToSet(r.nftSet.Name) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// we delete last IP from the set, that means we need to delete
|
|
||||||
// set itself and associated firewall rule too
|
|
||||||
m.rConn.FlushSet(r.nftSet)
|
|
||||||
m.rConn.DelSet(r.nftSet)
|
|
||||||
m.ipsetStore.deleteIpset(r.nftSet.Name)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// createDefaultAllowRules creates default allow rules for the input and output chains
|
|
||||||
func (m *AclManager) createDefaultAllowRules() error {
|
|
||||||
expIn := []expr.Any{
|
|
||||||
&expr.Verdict{
|
|
||||||
Kind: expr.VerdictAccept,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
_ = m.rConn.InsertRule(&nftables.Rule{
|
|
||||||
Table: m.workTable,
|
|
||||||
Chain: m.chainInputRules,
|
|
||||||
Position: 0,
|
|
||||||
Exprs: expIn,
|
|
||||||
})
|
|
||||||
|
|
||||||
if err := m.rConn.Flush(); err != nil {
|
|
||||||
return fmt.Errorf(flushError, err)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Flush rule/chain/set operations from the buffer
|
|
||||||
//
|
|
||||||
// Method also get all rules after flush and refreshes handle values in the rulesets
|
|
||||||
func (m *AclManager) Flush() error {
|
|
||||||
if err := m.flushWithBackoff(); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.refreshRuleHandles(m.chainInputRules, false); err != nil {
|
|
||||||
log.Errorf("failed to refresh rule handles ipv4 input chain: %v", err)
|
|
||||||
}
|
|
||||||
if err := m.refreshRuleHandles(m.chainPrerouting, true); err != nil {
|
|
||||||
log.Errorf("failed to refresh rule handles prerouting chain: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) addIOFiltering(
|
|
||||||
ip net.IP,
|
|
||||||
proto firewall.Protocol,
|
|
||||||
sPort *firewall.Port,
|
|
||||||
dPort *firewall.Port,
|
|
||||||
action firewall.Action,
|
|
||||||
ipset *nftables.Set,
|
|
||||||
) (*Rule, error) {
|
|
||||||
ruleId := generatePeerRuleId(ip, proto, sPort, dPort, action, ipset)
|
|
||||||
if r, ok := m.rules[ruleId]; ok {
|
|
||||||
return &Rule{
|
|
||||||
nftRule: r.nftRule,
|
|
||||||
mangleRule: r.mangleRule,
|
|
||||||
nftSet: r.nftSet,
|
|
||||||
ruleID: r.ruleID,
|
|
||||||
ip: ip,
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var expressions []expr.Any
|
|
||||||
|
|
||||||
if proto != firewall.ProtocolALL {
|
|
||||||
expressions = append(expressions, &expr.Payload{
|
|
||||||
DestRegister: 1,
|
|
||||||
Base: expr.PayloadBaseNetworkHeader,
|
|
||||||
Offset: m.af.protoOffset,
|
|
||||||
Len: uint32(1),
|
|
||||||
})
|
|
||||||
|
|
||||||
protoData, err := m.af.protoNum(proto)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("convert protocol to number: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
expressions = append(expressions, &expr.Cmp{
|
|
||||||
Register: 1,
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Data: []byte{protoData},
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
rawIP := ipToBytes(ip, m.af)
|
|
||||||
// check if rawIP contains zeroed IPv4 0.0.0.0 value
|
|
||||||
// in that case not add IP match expression into the rule definition
|
|
||||||
if slices.ContainsFunc(rawIP, func(v byte) bool { return v != 0 }) {
|
|
||||||
expressions = append(expressions,
|
|
||||||
&expr.Payload{
|
|
||||||
DestRegister: 1,
|
|
||||||
Base: expr.PayloadBaseNetworkHeader,
|
|
||||||
Offset: m.af.srcAddrOffset,
|
|
||||||
Len: m.af.addrLen,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
// add individual IP for match if no ipset defined
|
|
||||||
if ipset == nil {
|
|
||||||
expressions = append(expressions,
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Register: 1,
|
|
||||||
Data: rawIP,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
} else {
|
|
||||||
expressions = append(expressions,
|
|
||||||
&expr.Lookup{
|
|
||||||
SourceRegister: 1,
|
|
||||||
SetName: ipset.Name,
|
|
||||||
SetID: ipset.ID,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
expressions = append(expressions, applyPort(sPort, true)...)
|
|
||||||
expressions = append(expressions, applyPort(dPort, false)...)
|
|
||||||
|
|
||||||
mainExpressions := slices.Clone(expressions)
|
|
||||||
|
|
||||||
switch action {
|
|
||||||
case firewall.ActionAccept:
|
|
||||||
mainExpressions = append(mainExpressions, &expr.Verdict{Kind: expr.VerdictAccept})
|
|
||||||
case firewall.ActionDrop:
|
|
||||||
mainExpressions = append(mainExpressions, &expr.Verdict{Kind: expr.VerdictDrop})
|
|
||||||
}
|
|
||||||
|
|
||||||
userData := []byte(ruleId)
|
|
||||||
|
|
||||||
chain := m.chainInputRules
|
|
||||||
rule := &nftables.Rule{
|
|
||||||
Table: m.workTable,
|
|
||||||
Chain: chain,
|
|
||||||
Exprs: mainExpressions,
|
|
||||||
UserData: userData,
|
|
||||||
}
|
|
||||||
|
|
||||||
// Insert DROP rules at the beginning, append ACCEPT rules at the end
|
|
||||||
var nftRule *nftables.Rule
|
|
||||||
if action == firewall.ActionDrop {
|
|
||||||
nftRule = m.rConn.InsertRule(rule)
|
|
||||||
} else {
|
|
||||||
nftRule = m.rConn.AddRule(rule)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.rConn.Flush(); err != nil {
|
|
||||||
return nil, fmt.Errorf("flush input rule %s: %v", ruleId, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ruleStruct := &Rule{
|
|
||||||
nftRule: nftRule,
|
|
||||||
// best effort mangle rule
|
|
||||||
mangleRule: m.createPreroutingRule(expressions, userData),
|
|
||||||
nftSet: ipset,
|
|
||||||
ruleID: ruleId,
|
|
||||||
ip: ip,
|
|
||||||
}
|
|
||||||
m.rules[ruleId] = ruleStruct
|
|
||||||
if ipset != nil {
|
|
||||||
m.ipsetStore.AddReferenceToIpset(ipset.Name)
|
|
||||||
}
|
|
||||||
|
|
||||||
return ruleStruct, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) createPreroutingRule(expressions []expr.Any, userData []byte) *nftables.Rule {
|
|
||||||
if m.chainPrerouting == nil {
|
|
||||||
log.Warn("prerouting chain is not created")
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
preroutingExprs := slices.Clone(expressions)
|
|
||||||
|
|
||||||
// interface
|
|
||||||
preroutingExprs = append([]expr.Any{
|
|
||||||
&expr.Meta{
|
|
||||||
Key: expr.MetaKeyIIFNAME,
|
|
||||||
Register: 1,
|
|
||||||
},
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Register: 1,
|
|
||||||
Data: ifname(m.wgIface.Name()),
|
|
||||||
},
|
|
||||||
}, preroutingExprs...)
|
|
||||||
|
|
||||||
// local destination and mark
|
|
||||||
preroutingExprs = append(preroutingExprs,
|
|
||||||
&expr.Fib{
|
|
||||||
Register: 1,
|
|
||||||
ResultADDRTYPE: true,
|
|
||||||
FlagDADDR: true,
|
|
||||||
},
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Register: 1,
|
|
||||||
Data: binaryutil.NativeEndian.PutUint32(unix.RTN_LOCAL),
|
|
||||||
},
|
|
||||||
|
|
||||||
&expr.Immediate{
|
|
||||||
Register: 1,
|
|
||||||
Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkRedirected),
|
|
||||||
},
|
|
||||||
&expr.Meta{
|
|
||||||
Key: expr.MetaKeyMARK,
|
|
||||||
Register: 1,
|
|
||||||
SourceRegister: true,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
nfRule := m.rConn.AddRule(&nftables.Rule{
|
|
||||||
Table: m.workTable,
|
|
||||||
Chain: m.chainPrerouting,
|
|
||||||
Exprs: preroutingExprs,
|
|
||||||
UserData: userData,
|
|
||||||
})
|
|
||||||
|
|
||||||
if err := m.rConn.Flush(); err != nil {
|
|
||||||
log.Errorf("failed to flush mangle rule %s: %v", string(userData), err)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return nfRule
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) createDefaultChains() (err error) {
|
|
||||||
// chainNameInputRules
|
|
||||||
chain := m.createChain(chainNameInputRules)
|
|
||||||
err = m.rConn.Flush()
|
|
||||||
if err != nil {
|
|
||||||
log.Debugf("failed to create chain (%s): %s", chain.Name, err)
|
|
||||||
return fmt.Errorf(flushError, err)
|
|
||||||
}
|
|
||||||
m.chainInputRules = chain
|
|
||||||
|
|
||||||
// netbird-acl-input-filter
|
|
||||||
// type filter hook input priority filter; policy accept;
|
|
||||||
chain = m.createFilterChainWithHook(chainNameInputFilter, nftables.ChainHookInput)
|
|
||||||
m.addJumpRule(chain, m.chainInputRules.Name, expr.MetaKeyIIFNAME) // to netbird-acl-input-rules
|
|
||||||
m.addDropExpressions(chain, expr.MetaKeyIIFNAME)
|
|
||||||
err = m.rConn.Flush()
|
|
||||||
if err != nil {
|
|
||||||
log.Debugf("failed to create chain (%s): %s", chain.Name, err)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
// netbird-acl-forward-filter
|
|
||||||
chainFwFilter := m.createFilterChainWithHook(chainNameForwardFilter, nftables.ChainHookForward)
|
|
||||||
m.addJumpRulesToRtForward(chainFwFilter) // to netbird-rt-fwd
|
|
||||||
m.addDropExpressions(chainFwFilter, expr.MetaKeyIIFNAME)
|
|
||||||
|
|
||||||
err = m.rConn.Flush()
|
|
||||||
if err != nil {
|
|
||||||
log.Debugf("failed to create chain (%s): %s", chainNameForwardFilter, err)
|
|
||||||
return fmt.Errorf(flushError, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.allowRedirectedTraffic(chainFwFilter); err != nil {
|
|
||||||
log.Errorf("failed to allow redirected traffic: %s", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Makes redirected traffic originally destined for the host itself (now subject to the forward filter)
|
|
||||||
// go through the input filter as well. This will enable e.g. Docker services to keep working by accessing the
|
|
||||||
// netbird peer IP.
|
|
||||||
func (m *AclManager) allowRedirectedTraffic(chainFwFilter *nftables.Chain) error {
|
|
||||||
// Chain is created by route manager
|
|
||||||
// TODO: move creation to a common place
|
|
||||||
m.chainPrerouting = &nftables.Chain{
|
|
||||||
Name: chainNameManglePrerouting,
|
|
||||||
Table: m.workTable,
|
|
||||||
Type: nftables.ChainTypeFilter,
|
|
||||||
Hooknum: nftables.ChainHookPrerouting,
|
|
||||||
Priority: nftables.ChainPriorityMangle,
|
|
||||||
}
|
|
||||||
|
|
||||||
m.addFwmarkToForward(chainFwFilter)
|
|
||||||
|
|
||||||
if err := m.rConn.Flush(); err != nil {
|
|
||||||
return fmt.Errorf(flushError, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) addFwmarkToForward(chainFwFilter *nftables.Chain) {
|
|
||||||
m.rConn.InsertRule(&nftables.Rule{
|
|
||||||
Table: m.workTable,
|
|
||||||
Chain: chainFwFilter,
|
|
||||||
Exprs: []expr.Any{
|
|
||||||
&expr.Meta{
|
|
||||||
Key: expr.MetaKeyMARK,
|
|
||||||
Register: 1,
|
|
||||||
},
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Register: 1,
|
|
||||||
Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkRedirected),
|
|
||||||
},
|
|
||||||
&expr.Verdict{
|
|
||||||
Kind: expr.VerdictAccept,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) addJumpRulesToRtForward(chainFwFilter *nftables.Chain) {
|
|
||||||
expressions := []expr.Any{
|
|
||||||
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Register: 1,
|
|
||||||
Data: ifname(m.wgIface.Name()),
|
|
||||||
},
|
|
||||||
&expr.Verdict{
|
|
||||||
Kind: expr.VerdictJump,
|
|
||||||
Chain: m.routingFwChainName,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
_ = m.rConn.AddRule(&nftables.Rule{
|
|
||||||
Table: m.workTable,
|
|
||||||
Chain: chainFwFilter,
|
|
||||||
Exprs: expressions,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) createChain(name string) *nftables.Chain {
|
|
||||||
chain := &nftables.Chain{
|
|
||||||
Name: name,
|
|
||||||
Table: m.workTable,
|
|
||||||
}
|
|
||||||
|
|
||||||
chain = m.rConn.AddChain(chain)
|
|
||||||
|
|
||||||
insertReturnTrafficRule(m.rConn, m.workTable, chain)
|
|
||||||
|
|
||||||
return chain
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) createFilterChainWithHook(name string, hookNum *nftables.ChainHook) *nftables.Chain {
|
|
||||||
polAccept := nftables.ChainPolicyAccept
|
|
||||||
chain := &nftables.Chain{
|
|
||||||
Name: name,
|
|
||||||
Table: m.workTable,
|
|
||||||
Hooknum: hookNum,
|
|
||||||
Priority: nftables.ChainPriorityFilter,
|
|
||||||
Type: nftables.ChainTypeFilter,
|
|
||||||
Policy: &polAccept,
|
|
||||||
}
|
|
||||||
|
|
||||||
return m.rConn.AddChain(chain)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) addDropExpressions(chain *nftables.Chain, ifaceKey expr.MetaKey) []expr.Any {
|
|
||||||
expressions := []expr.Any{
|
|
||||||
&expr.Meta{Key: ifaceKey, Register: 1},
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Register: 1,
|
|
||||||
Data: ifname(m.wgIface.Name()),
|
|
||||||
},
|
|
||||||
&expr.Verdict{Kind: expr.VerdictDrop},
|
|
||||||
}
|
|
||||||
_ = m.rConn.AddRule(&nftables.Rule{
|
|
||||||
Table: m.workTable,
|
|
||||||
Chain: chain,
|
|
||||||
Exprs: expressions,
|
|
||||||
})
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) addJumpRule(chain *nftables.Chain, to string, ifaceKey expr.MetaKey) {
|
|
||||||
expressions := []expr.Any{
|
|
||||||
&expr.Meta{Key: ifaceKey, Register: 1},
|
|
||||||
&expr.Cmp{
|
|
||||||
Op: expr.CmpOpEq,
|
|
||||||
Register: 1,
|
|
||||||
Data: ifname(m.wgIface.Name()),
|
|
||||||
},
|
|
||||||
&expr.Verdict{
|
|
||||||
Kind: expr.VerdictJump,
|
|
||||||
Chain: to,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
_ = m.rConn.AddRule(&nftables.Rule{
|
|
||||||
Table: chain.Table,
|
|
||||||
Chain: chain,
|
|
||||||
Exprs: expressions,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) addIpToSet(ipsetName string, ip net.IP) (*nftables.Set, error) {
|
|
||||||
ipset, err := m.rConn.GetSetByName(m.workTable, ipsetName)
|
|
||||||
rawIP := ipToBytes(ip, m.af)
|
|
||||||
if err != nil {
|
|
||||||
if ipset, err = m.createSet(m.workTable, ipsetName); err != nil {
|
|
||||||
return nil, fmt.Errorf("get set name: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.ipsetStore.newIpset(ipset.Name)
|
|
||||||
}
|
|
||||||
|
|
||||||
if m.ipsetStore.IsIpInSet(ipset.Name, ip) {
|
|
||||||
return ipset, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.sConn.SetAddElements(ipset, []nftables.SetElement{{Key: rawIP}}); err != nil {
|
|
||||||
return nil, fmt.Errorf("add set element for the first time: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.ipsetStore.AddIpToSet(ipset.Name, ip)
|
|
||||||
|
|
||||||
if err := m.sConn.Flush(); err != nil {
|
|
||||||
return nil, fmt.Errorf("flush add elements: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return ipset, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// createSet in given table by name
|
|
||||||
func (m *AclManager) createSet(table *nftables.Table, name string) (*nftables.Set, error) {
|
|
||||||
ipset := &nftables.Set{
|
|
||||||
Name: name,
|
|
||||||
Table: table,
|
|
||||||
Dynamic: true,
|
|
||||||
KeyType: m.af.setKeyType,
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.rConn.AddSet(ipset, nil); err != nil {
|
|
||||||
return nil, fmt.Errorf("create set: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.rConn.Flush(); err != nil {
|
|
||||||
return nil, fmt.Errorf("flush created set: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return ipset, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) flushWithBackoff() (err error) {
|
|
||||||
backoff := 4
|
|
||||||
backoffTime := 1000 * time.Millisecond
|
|
||||||
for i := 0; ; i++ {
|
|
||||||
err = m.rConn.Flush()
|
|
||||||
if err != nil {
|
|
||||||
log.Debugf("failed to flush nftables: %v", err)
|
|
||||||
if !strings.Contains(err.Error(), "busy") {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
log.Error("failed to flush nftables, retrying...")
|
|
||||||
if i == backoff-1 {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
time.Sleep(backoffTime)
|
|
||||||
backoffTime *= 2
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
break
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *AclManager) refreshRuleHandles(chain *nftables.Chain, mangle bool) error {
|
|
||||||
if m.workTable == nil || chain == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
list, err := m.rConn.GetRules(m.workTable, chain)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, rule := range list {
|
|
||||||
if len(rule.UserData) == 0 {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
split := bytes.Split(rule.UserData, []byte(" "))
|
|
||||||
r, ok := m.rules[string(split[0])]
|
|
||||||
if ok {
|
|
||||||
if mangle {
|
|
||||||
*r.mangleRule = *rule
|
|
||||||
} else {
|
|
||||||
*r.nftRule = *rule
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func generatePeerRuleId(ip net.IP, proto firewall.Protocol, sPort *firewall.Port, dPort *firewall.Port, action firewall.Action, ipset *nftables.Set) string {
|
|
||||||
rulesetID := ":" + string(proto) + ":"
|
|
||||||
if sPort != nil {
|
|
||||||
rulesetID += sPort.String()
|
|
||||||
}
|
|
||||||
rulesetID += ":"
|
|
||||||
if dPort != nil {
|
|
||||||
rulesetID += dPort.String()
|
|
||||||
}
|
|
||||||
rulesetID += ":"
|
|
||||||
rulesetID += strconv.Itoa(int(action))
|
|
||||||
if ipset == nil {
|
|
||||||
return "ip:" + ip.String() + rulesetID
|
|
||||||
}
|
|
||||||
return "set:" + ipset.Name + rulesetID
|
|
||||||
}
|
|
||||||
|
|
||||||
func ifname(n string) []byte {
|
|
||||||
b := make([]byte, 16)
|
|
||||||
copy(b, n+"\x00")
|
|
||||||
return b
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
// ipToBytes converts net.IP to the correct byte length for the address family.
|
|
||||||
func ipToBytes(ip net.IP, af addrFamily) []byte {
|
|
||||||
if af.addrLen == 4 {
|
|
||||||
return ip.To4()
|
|
||||||
}
|
|
||||||
return ip.To16()
|
|
||||||
}
|
|
||||||
|
|
||||||
@@ -0,0 +1,880 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/coreos/go-iptables/iptables"
|
||||||
|
"github.com/google/nftables"
|
||||||
|
"github.com/google/nftables/binaryutil"
|
||||||
|
"github.com/google/nftables/expr"
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
"github.com/netbirdio/netbird/client/firewall/firewalld"
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
nbnet "github.com/netbirdio/netbird/client/net"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (r *family) createContainers() error {
|
||||||
|
r.chains[chainNameRoutingFw] = r.conn.AddChain(&nftables.Chain{
|
||||||
|
Name: chainNameRoutingFw,
|
||||||
|
Table: r.workTable,
|
||||||
|
})
|
||||||
|
|
||||||
|
prio := *nftables.ChainPriorityNATSource - 1
|
||||||
|
r.chains[chainNameRoutingNat] = r.conn.AddChain(&nftables.Chain{
|
||||||
|
Name: chainNameRoutingNat,
|
||||||
|
Table: r.workTable,
|
||||||
|
Hooknum: nftables.ChainHookPostrouting,
|
||||||
|
Priority: &prio,
|
||||||
|
Type: nftables.ChainTypeNAT,
|
||||||
|
})
|
||||||
|
|
||||||
|
r.chains[chainNameRoutingRdr] = r.conn.AddChain(&nftables.Chain{
|
||||||
|
Name: chainNameRoutingRdr,
|
||||||
|
Table: r.workTable,
|
||||||
|
Hooknum: nftables.ChainHookPrerouting,
|
||||||
|
Priority: nftables.ChainPriorityNATDest,
|
||||||
|
Type: nftables.ChainTypeNAT,
|
||||||
|
})
|
||||||
|
|
||||||
|
r.chains[chainNameManglePostrouting] = r.conn.AddChain(&nftables.Chain{
|
||||||
|
Name: chainNameManglePostrouting,
|
||||||
|
Table: r.workTable,
|
||||||
|
Hooknum: nftables.ChainHookPostrouting,
|
||||||
|
Priority: nftables.ChainPriorityMangle,
|
||||||
|
Type: nftables.ChainTypeFilter,
|
||||||
|
})
|
||||||
|
|
||||||
|
r.chains[chainNameManglePrerouting] = r.conn.AddChain(&nftables.Chain{
|
||||||
|
Name: chainNameManglePrerouting,
|
||||||
|
Table: r.workTable,
|
||||||
|
Hooknum: nftables.ChainHookPrerouting,
|
||||||
|
Priority: nftables.ChainPriorityMangle,
|
||||||
|
Type: nftables.ChainTypeFilter,
|
||||||
|
})
|
||||||
|
|
||||||
|
r.chains[chainNameMangleForward] = r.conn.AddChain(&nftables.Chain{
|
||||||
|
Name: chainNameMangleForward,
|
||||||
|
Table: r.workTable,
|
||||||
|
Hooknum: nftables.ChainHookForward,
|
||||||
|
Priority: nftables.ChainPriorityMangle,
|
||||||
|
Type: nftables.ChainTypeFilter,
|
||||||
|
})
|
||||||
|
|
||||||
|
insertReturnTrafficRule(r.conn, r.workTable, r.chains[chainNameRoutingFw])
|
||||||
|
|
||||||
|
r.addPostroutingRules()
|
||||||
|
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf("initialize tables: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.addMSSClampingRules(); err != nil {
|
||||||
|
log.Errorf("failed to add MSS clamping rules: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Kernel routing opens both INPUT and FORWARD.
|
||||||
|
if err := r.openInterface(true); err != nil {
|
||||||
|
log.Errorf("failed to open interface in foreign chains: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := firewalld.TrustInterface(r.wgIface.Name()); err != nil {
|
||||||
|
log.Warnf("failed to trust interface in firewalld: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.refreshRulesMap(); err != nil {
|
||||||
|
log.Errorf("failed to refresh rules: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// setupDataPlaneMark configures the fwmark for the data plane
|
||||||
|
func (r *family) setupDataPlaneMark() error {
|
||||||
|
if r.chains[chainNameManglePrerouting] == nil || r.chains[chainNameManglePostrouting] == nil {
|
||||||
|
return errors.New("no mangle chains found")
|
||||||
|
}
|
||||||
|
|
||||||
|
ctNew := getCtNewExprs()
|
||||||
|
preExprs := []expr.Any{
|
||||||
|
&expr.Meta{
|
||||||
|
Key: expr.MetaKeyIIFNAME,
|
||||||
|
Register: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname(r.wgIface.Name()),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
preExprs = append(preExprs, ctNew...)
|
||||||
|
preExprs = append(preExprs,
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 1,
|
||||||
|
Data: binaryutil.NativeEndian.PutUint32(nbnet.DataPlaneMarkIn),
|
||||||
|
},
|
||||||
|
&expr.Ct{
|
||||||
|
Key: expr.CtKeyMARK,
|
||||||
|
Register: 1,
|
||||||
|
SourceRegister: true,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
preNftRule := &nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chains[chainNameManglePrerouting],
|
||||||
|
Exprs: preExprs,
|
||||||
|
}
|
||||||
|
r.conn.AddRule(preNftRule)
|
||||||
|
|
||||||
|
postExprs := []expr.Any{
|
||||||
|
&expr.Meta{
|
||||||
|
Key: expr.MetaKeyOIFNAME,
|
||||||
|
Register: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname(r.wgIface.Name()),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
postExprs = append(postExprs, ctNew...)
|
||||||
|
postExprs = append(postExprs,
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 1,
|
||||||
|
Data: binaryutil.NativeEndian.PutUint32(nbnet.DataPlaneMarkOut),
|
||||||
|
},
|
||||||
|
&expr.Ct{
|
||||||
|
Key: expr.CtKeyMARK,
|
||||||
|
Register: 1,
|
||||||
|
SourceRegister: true,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
postNftRule := &nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chains[chainNameManglePostrouting],
|
||||||
|
Exprs: postExprs,
|
||||||
|
}
|
||||||
|
r.conn.AddRule(postNftRule)
|
||||||
|
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf("flush: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// openInterface adds passthrough accept rules for the NetBird interface to the
|
||||||
|
// kernel's filter table and external chains so they don't drop our traffic.
|
||||||
|
// includeForward also opens the FORWARD chains (kernel routing); when false only
|
||||||
|
// INPUT is opened, which is all the userspace router needs since it never
|
||||||
|
// forwards in the kernel.
|
||||||
|
func (r *family) openInterface(includeForward bool) error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
|
||||||
|
if err := r.acceptFilterTableRules(includeForward); err != nil {
|
||||||
|
merr = multierror.Append(merr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.acceptExternalChainsRules(includeForward); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("add accept rules to external chains: %w", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) acceptFilterTableRules(includeForward bool) error {
|
||||||
|
if r.filterTable == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
fw := "iptables"
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
log.Debugf("Used %s to add accept input/forward rules", fw)
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Try iptables first and fallback to nftables if iptables is not available.
|
||||||
|
// Use the correct protocol (iptables vs ip6tables) for the address family.
|
||||||
|
ipt, err := iptables.NewWithProtocol(r.iptablesProto())
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("Will use nftables to manipulate the filter table because iptables is not available: %v", err)
|
||||||
|
|
||||||
|
fw = "nftables"
|
||||||
|
return r.acceptFilterRulesNftables(r.filterTable, includeForward)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.acceptFilterRulesIptables(ipt, includeForward); err != nil {
|
||||||
|
log.Warnf("iptables failed (table may be incompatible), falling back to nftables: %v", err)
|
||||||
|
fw = "nftables"
|
||||||
|
return r.acceptFilterRulesNftables(r.filterTable, includeForward)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) acceptFilterRulesIptables(ipt *iptables.IPTables, includeForward bool) error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
|
||||||
|
if includeForward {
|
||||||
|
for _, rule := range r.getAcceptForwardRules() {
|
||||||
|
if err := ipt.Insert("filter", chainNameForward, 1, rule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("add iptables forward rule: %v", err))
|
||||||
|
} else {
|
||||||
|
log.Debugf("added iptables forward rule: %v", rule)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
inputRule := r.getAcceptInputRule()
|
||||||
|
if err := ipt.Insert("filter", chainNameInput, 1, inputRule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("add iptables input rule: %v", err))
|
||||||
|
} else {
|
||||||
|
log.Debugf("added iptables input rule: %v", inputRule)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) getAcceptForwardRules() [][]string {
|
||||||
|
intf := r.wgIface.Name()
|
||||||
|
return [][]string{
|
||||||
|
{"-i", intf, "-j", "ACCEPT"},
|
||||||
|
{"-o", intf, "-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) getAcceptInputRule() []string {
|
||||||
|
return []string{"-i", r.wgIface.Name(), "-j", "ACCEPT"}
|
||||||
|
}
|
||||||
|
|
||||||
|
// acceptFilterRulesNftables adds accept rules to the ip filter table using nftables.
|
||||||
|
// This is used when iptables is not available.
|
||||||
|
func (r *family) acceptFilterRulesNftables(table *nftables.Table, includeForward bool) error {
|
||||||
|
intf := ifname(r.wgIface.Name())
|
||||||
|
|
||||||
|
if includeForward {
|
||||||
|
forwardChain := &nftables.Chain{
|
||||||
|
Name: chainNameForward,
|
||||||
|
Table: table,
|
||||||
|
Type: nftables.ChainTypeFilter,
|
||||||
|
Hooknum: nftables.ChainHookForward,
|
||||||
|
Priority: nftables.ChainPriorityFilter,
|
||||||
|
}
|
||||||
|
r.insertForwardAcceptRules(forwardChain, intf)
|
||||||
|
}
|
||||||
|
|
||||||
|
inputChain := &nftables.Chain{
|
||||||
|
Name: chainNameInput,
|
||||||
|
Table: table,
|
||||||
|
Type: nftables.ChainTypeFilter,
|
||||||
|
Hooknum: nftables.ChainHookInput,
|
||||||
|
Priority: nftables.ChainPriorityFilter,
|
||||||
|
}
|
||||||
|
r.insertInputAcceptRule(inputChain, intf)
|
||||||
|
|
||||||
|
return r.conn.Flush()
|
||||||
|
}
|
||||||
|
|
||||||
|
// acceptExternalChainsRules adds accept rules to external chains (non-netbird, non-iptables tables).
|
||||||
|
// It dynamically finds chains at call time to handle chains that may have been created after startup.
|
||||||
|
func (r *family) acceptExternalChainsRules(includeForward bool) error {
|
||||||
|
chains := r.findExternalChains()
|
||||||
|
if len(chains) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
intf := ifname(r.wgIface.Name())
|
||||||
|
for _, chain := range chains {
|
||||||
|
r.applyExternalChainAccept(chain, intf, includeForward)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf("flush external chain rules: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) applyExternalChainAccept(chain *nftables.Chain, intf []byte, includeForward bool) {
|
||||||
|
if chain.Hooknum == nil {
|
||||||
|
log.Debugf("skipping external chain %s/%s: hooknum is nil", chain.Table.Name, chain.Name)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugf("adding accept rules to external %s chain: %s %s/%s",
|
||||||
|
hookName(chain.Hooknum), familyName(chain.Table.Family), chain.Table.Name, chain.Name)
|
||||||
|
|
||||||
|
switch *chain.Hooknum {
|
||||||
|
case *nftables.ChainHookForward:
|
||||||
|
if includeForward {
|
||||||
|
r.insertForwardAcceptRules(chain, intf)
|
||||||
|
}
|
||||||
|
case *nftables.ChainHookInput:
|
||||||
|
r.insertInputAcceptRule(chain, intf)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) insertForwardAcceptRules(chain *nftables.Chain, intf []byte) {
|
||||||
|
existing, err := r.existingNetbirdRulesInChain(chain)
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("skip forward accept rules in %s/%s: %v", chain.Table.Name, chain.Name, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.insertForwardIifRule(chain, intf, existing)
|
||||||
|
r.insertForwardOifEstablishedRule(chain, intf, existing)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) insertForwardIifRule(chain *nftables.Chain, intf []byte, existing map[string]bool) {
|
||||||
|
if existing[userDataAcceptForwardRuleIif] {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.conn.InsertRule(&nftables.Rule{
|
||||||
|
Table: chain.Table,
|
||||||
|
Chain: chain,
|
||||||
|
Exprs: []expr.Any{
|
||||||
|
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: intf},
|
||||||
|
&expr.Counter{},
|
||||||
|
&expr.Verdict{Kind: expr.VerdictAccept},
|
||||||
|
},
|
||||||
|
UserData: []byte(userDataAcceptForwardRuleIif),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) insertForwardOifEstablishedRule(chain *nftables.Chain, intf []byte, existing map[string]bool) {
|
||||||
|
if existing[userDataAcceptForwardRuleOif] {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
exprs := []expr.Any{
|
||||||
|
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: intf},
|
||||||
|
}
|
||||||
|
r.conn.InsertRule(&nftables.Rule{
|
||||||
|
Table: chain.Table,
|
||||||
|
Chain: chain,
|
||||||
|
Exprs: append(exprs, getEstablishedExprs(2)...),
|
||||||
|
UserData: []byte(userDataAcceptForwardRuleOif),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) insertInputAcceptRule(chain *nftables.Chain, intf []byte) {
|
||||||
|
existing, err := r.existingNetbirdRulesInChain(chain)
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("skip input accept rule in %s/%s: %v", chain.Table.Name, chain.Name, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if existing[userDataAcceptInputRule] {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.conn.InsertRule(&nftables.Rule{
|
||||||
|
Table: chain.Table,
|
||||||
|
Chain: chain,
|
||||||
|
Exprs: []expr.Any{
|
||||||
|
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: intf},
|
||||||
|
&expr.Counter{},
|
||||||
|
&expr.Verdict{Kind: expr.VerdictAccept},
|
||||||
|
},
|
||||||
|
UserData: []byte(userDataAcceptInputRule),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// existingNetbirdRulesInChain returns the set of netbird-owned UserData tags present in a chain; callers must bail on error since InsertRule is additive.
|
||||||
|
func (r *family) existingNetbirdRulesInChain(chain *nftables.Chain) (map[string]bool, error) {
|
||||||
|
rules, err := r.conn.GetRules(chain.Table, chain)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("list rules: %w", err)
|
||||||
|
}
|
||||||
|
present := map[string]bool{}
|
||||||
|
for _, rule := range rules {
|
||||||
|
if !isNetbirdAcceptRuleTag(rule.UserData) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
present[string(rule.UserData)] = true
|
||||||
|
}
|
||||||
|
return present, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func isNetbirdAcceptRuleTag(userData []byte) bool {
|
||||||
|
switch string(userData) {
|
||||||
|
case userDataAcceptForwardRuleIif,
|
||||||
|
userDataAcceptForwardRuleOif,
|
||||||
|
userDataAcceptInputRule:
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) removeAcceptFilterRules() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
|
||||||
|
if err := r.removeFilterTableRules(); err != nil {
|
||||||
|
merr = multierror.Append(merr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.removeExternalChainsRules(); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove external chain rules: %w", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) removeFilterTableRules() error {
|
||||||
|
if r.filterTable == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
ipt, err := iptables.NewWithProtocol(r.iptablesProto())
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("iptables not available, using nftables to remove filter rules: %v", err)
|
||||||
|
return r.removeAcceptRulesFromTable(r.filterTable)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.removeAcceptFilterRulesIptables(ipt); err != nil {
|
||||||
|
log.Debugf("iptables removal failed (table may be incompatible), falling back to nftables: %v", err)
|
||||||
|
return r.removeAcceptRulesFromTable(r.filterTable)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) removeAcceptRulesFromTable(table *nftables.Table) error {
|
||||||
|
chains, err := r.conn.ListChainsOfTableFamily(table.Family)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("list chains: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, chain := range chains {
|
||||||
|
if chain.Table.Name != table.Name {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if chain.Name != chainNameForward && chain.Name != chainNameInput {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.removeAcceptRulesFromChain(table, chain); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return r.conn.Flush()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) removeAcceptRulesFromChain(table *nftables.Table, chain *nftables.Chain) error {
|
||||||
|
rules, err := r.conn.GetRules(table, chain)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("get rules from %s/%s: %v", table.Name, chain.Name, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, rule := range rules {
|
||||||
|
if bytes.Equal(rule.UserData, []byte(userDataAcceptForwardRuleIif)) ||
|
||||||
|
bytes.Equal(rule.UserData, []byte(userDataAcceptForwardRuleOif)) ||
|
||||||
|
bytes.Equal(rule.UserData, []byte(userDataAcceptInputRule)) {
|
||||||
|
if err := r.conn.DelRule(rule); err != nil {
|
||||||
|
return fmt.Errorf("delete rule from %s/%s: %v", table.Name, chain.Name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// removeExternalChainsRules removes our accept rules from all external chains.
|
||||||
|
// This is deterministic - it scans for chains at removal time rather than relying on saved state,
|
||||||
|
// ensuring cleanup works even after a crash or if chains changed.
|
||||||
|
func (r *family) removeExternalChainsRules() error {
|
||||||
|
chains := r.findExternalChains()
|
||||||
|
if len(chains) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var merr *multierror.Error
|
||||||
|
for _, chain := range chains {
|
||||||
|
if err := r.removeAcceptRulesFromChain(chain.Table, chain); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove rules from external chain %s/%s: %w", chain.Table.Name, chain.Name, err))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("flush external chain %s/%s: %w", chain.Table.Name, chain.Name, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// findExternalChains scans for chains from non-netbird tables that have FORWARD or INPUT hooks.
|
||||||
|
// This is used both at startup (to know where to add rules) and at cleanup (to ensure deterministic removal).
|
||||||
|
func (r *family) findExternalChains() []*nftables.Chain {
|
||||||
|
var chains []*nftables.Chain
|
||||||
|
|
||||||
|
families := []nftables.TableFamily{r.af.tableFamily, nftables.TableFamilyINet}
|
||||||
|
|
||||||
|
for _, family := range families {
|
||||||
|
allChains, err := r.conn.ListChainsOfTableFamily(family)
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("list chains for family %d: %v", family, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, chain := range allChains {
|
||||||
|
if r.isExternalChain(chain) {
|
||||||
|
chains = append(chains, chain)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return chains
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) isExternalChain(chain *nftables.Chain) bool {
|
||||||
|
if r.workTable != nil && chain.Table.Name == r.workTable.Name {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Skip firewalld-owned chains. Firewalld creates its chains with the
|
||||||
|
// NFT_CHAIN_OWNER flag, so inserting rules into them returns EPERM.
|
||||||
|
// We delegate acceptance to firewalld by trusting the interface instead.
|
||||||
|
if chain.Table.Name == firewalldTableName {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Skip iptables/ip6tables-managed tables (adding nft-native rules breaks iptables-save compat)
|
||||||
|
if (chain.Table.Family == nftables.TableFamilyIPv4 || chain.Table.Family == nftables.TableFamilyIPv6) && isIptablesTable(chain.Table.Name) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if chain.Type != nftables.ChainTypeFilter {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if chain.Hooknum == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
return *chain.Hooknum == *nftables.ChainHookForward || *chain.Hooknum == *nftables.ChainHookInput
|
||||||
|
}
|
||||||
|
|
||||||
|
func isIptablesTable(name string) bool {
|
||||||
|
switch name {
|
||||||
|
case tableNameFilter, tableNat, tableMangle, tableRaw, tableSecurity:
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) removeAcceptFilterRulesIptables(ipt *iptables.IPTables) error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
|
||||||
|
for _, rule := range r.getAcceptForwardRules() {
|
||||||
|
if err := ipt.DeleteIfExists("filter", chainNameForward, rule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove iptables forward rule: %v", err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
inputRule := r.getAcceptInputRule()
|
||||||
|
if err := ipt.DeleteIfExists("filter", chainNameInput, inputRule...); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove iptables input rule: %v", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush rule/chain/set operations from the buffer
|
||||||
|
//
|
||||||
|
// Method also get all rules after flush and refreshes handle values in the rulesets
|
||||||
|
func (r *family) Flush() error {
|
||||||
|
if err := r.flushWithBackoff(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.refreshRuleHandles(r.chainInputRules, false); err != nil {
|
||||||
|
log.Errorf("failed to refresh rule handles ipv4 input chain: %v", err)
|
||||||
|
}
|
||||||
|
if err := r.refreshRuleHandles(r.chainPrerouting, true); err != nil {
|
||||||
|
log.Errorf("failed to refresh rule handles prerouting chain: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// queuePreroutingRule builds the prerouting mangle rule that marks
|
||||||
|
// redirected traffic and queues it on the connection without flushing,
|
||||||
|
// so the caller can commit it in the same transaction as the rule it
|
||||||
|
// pairs with. Returns nil when the prerouting chain is absent, in which
|
||||||
|
// case nothing is queued.
|
||||||
|
func (r *family) queuePreroutingRule(expressions []expr.Any, userData []byte) *nftables.Rule {
|
||||||
|
if r.chainPrerouting == nil {
|
||||||
|
log.Warn("prerouting chain is not created")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
preroutingExprs := slices.Clone(expressions)
|
||||||
|
|
||||||
|
// interface
|
||||||
|
preroutingExprs = append([]expr.Any{
|
||||||
|
&expr.Meta{
|
||||||
|
Key: expr.MetaKeyIIFNAME,
|
||||||
|
Register: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname(r.wgIface.Name()),
|
||||||
|
},
|
||||||
|
}, preroutingExprs...)
|
||||||
|
|
||||||
|
// local destination and mark
|
||||||
|
preroutingExprs = append(preroutingExprs,
|
||||||
|
&expr.Fib{
|
||||||
|
Register: 1,
|
||||||
|
ResultADDRTYPE: true,
|
||||||
|
FlagDADDR: true,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: binaryutil.NativeEndian.PutUint32(unix.RTN_LOCAL),
|
||||||
|
},
|
||||||
|
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 1,
|
||||||
|
Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkRedirected),
|
||||||
|
},
|
||||||
|
&expr.Meta{
|
||||||
|
Key: expr.MetaKeyMARK,
|
||||||
|
Register: 1,
|
||||||
|
SourceRegister: true,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
return r.conn.AddRule(&nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chainPrerouting,
|
||||||
|
Exprs: preroutingExprs,
|
||||||
|
UserData: userData,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) createDefaultChains() (err error) {
|
||||||
|
// chainNameInputRules
|
||||||
|
chain := r.createChain(chainNameInputRules)
|
||||||
|
err = r.conn.Flush()
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("failed to create chain (%s): %s", chain.Name, err)
|
||||||
|
return fmt.Errorf(flushError, err)
|
||||||
|
}
|
||||||
|
r.chainInputRules = chain
|
||||||
|
|
||||||
|
// netbird-acl-input-filter
|
||||||
|
// type filter hook input priority filter; policy accept;
|
||||||
|
chain = r.createFilterChainWithHook(chainNameInputFilter, nftables.ChainHookInput)
|
||||||
|
r.addJumpRule(chain, r.chainInputRules.Name, expr.MetaKeyIIFNAME) // to netbird-acl-input-rules
|
||||||
|
r.addDropExpressions(chain, expr.MetaKeyIIFNAME)
|
||||||
|
err = r.conn.Flush()
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("failed to create chain (%s): %s", chain.Name, err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// netbird-acl-forward-filter
|
||||||
|
chainFwFilter := r.createFilterChainWithHook(chainNameForwardFilter, nftables.ChainHookForward)
|
||||||
|
r.addJumpRulesToRtForward(chainFwFilter) // to netbird-rt-fwd
|
||||||
|
r.addDropExpressions(chainFwFilter, expr.MetaKeyIIFNAME)
|
||||||
|
|
||||||
|
err = r.conn.Flush()
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("failed to create chain (%s): %s", chainNameForwardFilter, err)
|
||||||
|
return fmt.Errorf(flushError, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.allowRedirectedTraffic(chainFwFilter); err != nil {
|
||||||
|
log.Errorf("failed to allow redirected traffic: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Makes redirected traffic originally destined for the host itself (now subject to the forward filter)
|
||||||
|
// go through the input filter as well. This will enable e.g. Docker services to keep working by accessing the
|
||||||
|
// netbird peer IP.
|
||||||
|
func (r *family) allowRedirectedTraffic(chainFwFilter *nftables.Chain) error {
|
||||||
|
r.chainPrerouting = r.chains[chainNameManglePrerouting]
|
||||||
|
|
||||||
|
r.addFwmarkToForward(chainFwFilter)
|
||||||
|
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf(flushError, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addFwmarkToForward(chainFwFilter *nftables.Chain) {
|
||||||
|
r.conn.InsertRule(&nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: chainFwFilter,
|
||||||
|
Exprs: []expr.Any{
|
||||||
|
&expr.Meta{
|
||||||
|
Key: expr.MetaKeyMARK,
|
||||||
|
Register: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkRedirected),
|
||||||
|
},
|
||||||
|
&expr.Verdict{
|
||||||
|
Kind: expr.VerdictAccept,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addJumpRulesToRtForward(chainFwFilter *nftables.Chain) {
|
||||||
|
expressions := []expr.Any{
|
||||||
|
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname(r.wgIface.Name()),
|
||||||
|
},
|
||||||
|
&expr.Verdict{
|
||||||
|
Kind: expr.VerdictJump,
|
||||||
|
Chain: r.routingFwChainName,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = r.conn.AddRule(&nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: chainFwFilter,
|
||||||
|
Exprs: expressions,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) createChain(name string) *nftables.Chain {
|
||||||
|
chain := &nftables.Chain{
|
||||||
|
Name: name,
|
||||||
|
Table: r.workTable,
|
||||||
|
}
|
||||||
|
|
||||||
|
chain = r.conn.AddChain(chain)
|
||||||
|
|
||||||
|
insertReturnTrafficRule(r.conn, r.workTable, chain)
|
||||||
|
|
||||||
|
return chain
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) createFilterChainWithHook(name string, hookNum *nftables.ChainHook) *nftables.Chain {
|
||||||
|
polAccept := nftables.ChainPolicyAccept
|
||||||
|
chain := &nftables.Chain{
|
||||||
|
Name: name,
|
||||||
|
Table: r.workTable,
|
||||||
|
Hooknum: hookNum,
|
||||||
|
Priority: nftables.ChainPriorityFilter,
|
||||||
|
Type: nftables.ChainTypeFilter,
|
||||||
|
Policy: &polAccept,
|
||||||
|
}
|
||||||
|
|
||||||
|
return r.conn.AddChain(chain)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addDropExpressions(chain *nftables.Chain, ifaceKey expr.MetaKey) []expr.Any {
|
||||||
|
expressions := []expr.Any{
|
||||||
|
&expr.Meta{Key: ifaceKey, Register: 1},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname(r.wgIface.Name()),
|
||||||
|
},
|
||||||
|
&expr.Verdict{Kind: expr.VerdictDrop},
|
||||||
|
}
|
||||||
|
_ = r.conn.AddRule(&nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: chain,
|
||||||
|
Exprs: expressions,
|
||||||
|
})
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addJumpRule(chain *nftables.Chain, to string, ifaceKey expr.MetaKey) {
|
||||||
|
expressions := []expr.Any{
|
||||||
|
&expr.Meta{Key: ifaceKey, Register: 1},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname(r.wgIface.Name()),
|
||||||
|
},
|
||||||
|
&expr.Verdict{
|
||||||
|
Kind: expr.VerdictJump,
|
||||||
|
Chain: to,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = r.conn.AddRule(&nftables.Rule{
|
||||||
|
Table: chain.Table,
|
||||||
|
Chain: chain,
|
||||||
|
Exprs: expressions,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) flushWithBackoff() (err error) {
|
||||||
|
backoff := 4
|
||||||
|
backoffTime := 1000 * time.Millisecond
|
||||||
|
for i := 0; ; i++ {
|
||||||
|
err = r.conn.Flush()
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("failed to flush nftables: %v", err)
|
||||||
|
if !strings.Contains(err.Error(), "busy") {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
log.Error("failed to flush nftables, retrying...")
|
||||||
|
if i == backoff-1 {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
time.Sleep(backoffTime)
|
||||||
|
backoffTime *= 2
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) refreshRuleHandles(chain *nftables.Chain, mangle bool) error {
|
||||||
|
if r.workTable == nil || chain == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
list, err := r.conn.GetRules(r.workTable, chain)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, rule := range list {
|
||||||
|
if len(rule.UserData) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
pr, ok := r.filters[firewall.RuleID(rule.UserData)]
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if mangle {
|
||||||
|
if pr.mangleRule != nil {
|
||||||
|
*pr.mangleRule = *rule
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
*pr.nftRule = *rule
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,573 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/google/nftables"
|
||||||
|
"github.com/google/nftables/binaryutil"
|
||||||
|
"github.com/google/nftables/expr"
|
||||||
|
"github.com/google/nftables/xt"
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (r *family) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) {
|
||||||
|
ruleID := rule.ID()
|
||||||
|
if _, exists := r.rules[ruleID+dnatSuffix]; exists {
|
||||||
|
return rule, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
protoNum, err := r.af.protoNum(rule.Protocol)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("convert protocol to number: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Request forwarding before queueing rules: addDnatRedirect/addDnatMasq
|
||||||
|
// buffer netlink messages on r.conn that the next caller's Flush would
|
||||||
|
// commit if we returned without flushing them ourselves.
|
||||||
|
if err := r.ipFwdState.RequestForwarding(r.isV6()); err != nil {
|
||||||
|
return nil, fmt.Errorf("enable forwarding: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.addDnatRedirect(rule, protoNum, ruleID); err != nil {
|
||||||
|
r.releaseForwarding()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.addDnatMasq(rule, protoNum, ruleID); err != nil {
|
||||||
|
r.releaseForwarding()
|
||||||
|
delete(r.rules, ruleID+dnatSuffix)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unlike iptables, there's no point in adding "out" rules in the forward chain here as our policy is ACCEPT.
|
||||||
|
// To overcome DROP policies in other chains, we'd have to add rules to the chains there.
|
||||||
|
// We also cannot just add "oif <iface> accept" there and filter in our own table as we don't know what is supposed to be allowed.
|
||||||
|
// TODO: find chains with drop policies and add rules there
|
||||||
|
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
r.releaseForwarding()
|
||||||
|
delete(r.rules, ruleID+dnatSuffix)
|
||||||
|
delete(r.rules, ruleID+snatSuffix)
|
||||||
|
return nil, fmt.Errorf("flush rules: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &rule, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addDnatRedirect(rule firewall.ForwardRule, protoNum uint8, ruleID firewall.RuleID) error {
|
||||||
|
dnatExprs := []expr.Any{
|
||||||
|
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpNeq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname(r.wgIface.Name()),
|
||||||
|
},
|
||||||
|
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: []byte{protoNum},
|
||||||
|
},
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 2,
|
||||||
|
Len: 2,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
portExprs, err := r.applyPort(&rule.DestinationPort, false)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("apply destination port: %w", err)
|
||||||
|
}
|
||||||
|
dnatExprs = append(dnatExprs, portExprs...)
|
||||||
|
|
||||||
|
// shifted translated port is not supported in nftables, so we hand this over to xtables
|
||||||
|
if rule.TranslatedPort.IsRange && len(rule.TranslatedPort.Values) == 2 {
|
||||||
|
if rule.TranslatedPort.Values[0] != rule.DestinationPort.Values[0] ||
|
||||||
|
rule.TranslatedPort.Values[1] != rule.DestinationPort.Values[1] {
|
||||||
|
return r.addXTablesRedirect(dnatExprs, ruleID, rule)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
additionalExprs, regProtoMin, regProtoMax, err := r.handleTranslatedPort(rule)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
dnatExprs = append(dnatExprs, additionalExprs...)
|
||||||
|
|
||||||
|
dnatExprs = append(dnatExprs,
|
||||||
|
&expr.NAT{
|
||||||
|
Type: expr.NATTypeDestNAT,
|
||||||
|
Family: uint32(r.af.tableFamily),
|
||||||
|
RegAddrMin: 1,
|
||||||
|
RegProtoMin: regProtoMin,
|
||||||
|
RegProtoMax: regProtoMax,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
dnatRule := &nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chains[chainNameRoutingRdr],
|
||||||
|
Exprs: dnatExprs,
|
||||||
|
UserData: []byte(ruleID + dnatSuffix),
|
||||||
|
}
|
||||||
|
r.conn.AddRule(dnatRule)
|
||||||
|
r.rules[ruleID+dnatSuffix] = dnatRule
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) handleTranslatedPort(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
|
||||||
|
switch {
|
||||||
|
case rule.TranslatedPort.IsRange && len(rule.TranslatedPort.Values) == 2:
|
||||||
|
return r.handlePortRange(rule)
|
||||||
|
case len(rule.TranslatedPort.Values) == 0:
|
||||||
|
return r.handleAddressOnly(rule)
|
||||||
|
case len(rule.TranslatedPort.Values) == 1:
|
||||||
|
return r.handleSinglePort(rule)
|
||||||
|
default:
|
||||||
|
return nil, 0, 0, fmt.Errorf("invalid translated port: %v", rule.TranslatedPort)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) handlePortRange(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
|
||||||
|
exprs := []expr.Any{
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 1,
|
||||||
|
Data: rule.TranslatedAddress.AsSlice(),
|
||||||
|
},
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 2,
|
||||||
|
Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[0]),
|
||||||
|
},
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 3,
|
||||||
|
Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[1]),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
return exprs, 2, 3, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) handleAddressOnly(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
|
||||||
|
exprs := []expr.Any{
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 1,
|
||||||
|
Data: rule.TranslatedAddress.AsSlice(),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
return exprs, 0, 0, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) handleSinglePort(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) {
|
||||||
|
exprs := []expr.Any{
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 1,
|
||||||
|
Data: rule.TranslatedAddress.AsSlice(),
|
||||||
|
},
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 2,
|
||||||
|
Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[0]),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
return exprs, 2, 0, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addXTablesRedirect(dnatExprs []expr.Any, ruleID firewall.RuleID, rule firewall.ForwardRule) error {
|
||||||
|
dnatExprs = append(dnatExprs,
|
||||||
|
&expr.Counter{},
|
||||||
|
&expr.Target{
|
||||||
|
Name: "DNAT",
|
||||||
|
Rev: 2,
|
||||||
|
Info: &xt.NatRange2{
|
||||||
|
NatRange: xt.NatRange{
|
||||||
|
Flags: uint(xt.NatRangeMapIPs | xt.NatRangeProtoSpecified | xt.NatRangeProtoOffset),
|
||||||
|
MinIP: rule.TranslatedAddress.AsSlice(),
|
||||||
|
MaxIP: rule.TranslatedAddress.AsSlice(),
|
||||||
|
MinPort: rule.TranslatedPort.Values[0],
|
||||||
|
MaxPort: rule.TranslatedPort.Values[1],
|
||||||
|
},
|
||||||
|
BasePort: rule.DestinationPort.Values[0],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
natTable := &nftables.Table{
|
||||||
|
Name: tableNat,
|
||||||
|
Family: r.af.tableFamily,
|
||||||
|
}
|
||||||
|
dnatRule := &nftables.Rule{
|
||||||
|
Table: natTable,
|
||||||
|
Chain: &nftables.Chain{
|
||||||
|
Name: chainNameNatPrerouting,
|
||||||
|
Table: natTable,
|
||||||
|
Type: nftables.ChainTypeNAT,
|
||||||
|
Hooknum: nftables.ChainHookPrerouting,
|
||||||
|
Priority: nftables.ChainPriorityNATDest,
|
||||||
|
},
|
||||||
|
Exprs: dnatExprs,
|
||||||
|
UserData: []byte(ruleID + dnatSuffix),
|
||||||
|
}
|
||||||
|
r.conn.AddRule(dnatRule)
|
||||||
|
r.rules[ruleID+dnatSuffix] = dnatRule
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addDnatMasq(rule firewall.ForwardRule, protoNum uint8, ruleID firewall.RuleID) error {
|
||||||
|
portExprs, err := r.applyPort(&rule.TranslatedPort, false)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("apply translated port: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
masqExprs := []expr.Any{
|
||||||
|
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname(r.wgIface.Name()),
|
||||||
|
},
|
||||||
|
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: []byte{protoNum},
|
||||||
|
},
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: r.af.dstAddrOffset,
|
||||||
|
Len: r.af.addrLen,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: rule.TranslatedAddress.AsSlice(),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
masqExprs = append(masqExprs, portExprs...)
|
||||||
|
masqExprs = append(masqExprs, &expr.Masq{})
|
||||||
|
|
||||||
|
masqRule := &nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chains[chainNameRoutingNat],
|
||||||
|
Exprs: masqExprs,
|
||||||
|
UserData: []byte(ruleID + snatSuffix),
|
||||||
|
}
|
||||||
|
r.conn.AddRule(masqRule)
|
||||||
|
r.rules[ruleID+snatSuffix] = masqRule
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) DeleteDNATRule(rule firewall.Rule) error {
|
||||||
|
ruleID := rule.ID()
|
||||||
|
|
||||||
|
if err := r.refreshRulesMap(); err != nil {
|
||||||
|
return fmt.Errorf(refreshRulesMapError, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var merr *multierror.Error
|
||||||
|
var needsFlush bool
|
||||||
|
var found bool
|
||||||
|
|
||||||
|
if dnatRule, exists := r.rules[ruleID+dnatSuffix]; exists {
|
||||||
|
found = true
|
||||||
|
if dnatRule.Handle == 0 {
|
||||||
|
log.Warnf("dnat rule %s has no handle, removing stale entry", ruleID+dnatSuffix)
|
||||||
|
delete(r.rules, ruleID+dnatSuffix)
|
||||||
|
} else if err := r.conn.DelRule(dnatRule); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete dnat rule: %w", err))
|
||||||
|
} else {
|
||||||
|
needsFlush = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if masqRule, exists := r.rules[ruleID+snatSuffix]; exists {
|
||||||
|
found = true
|
||||||
|
if masqRule.Handle == 0 {
|
||||||
|
log.Warnf("snat rule %s has no handle, removing stale entry", ruleID+snatSuffix)
|
||||||
|
delete(r.rules, ruleID+snatSuffix)
|
||||||
|
} else if err := r.conn.DelRule(masqRule); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete snat rule: %w", err))
|
||||||
|
} else {
|
||||||
|
needsFlush = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if needsFlush {
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf(flushError, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if merr != nil {
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
delete(r.rules, ruleID+dnatSuffix)
|
||||||
|
delete(r.rules, ruleID+snatSuffix)
|
||||||
|
|
||||||
|
// Release once, only if the rule was present and removed.
|
||||||
|
if found {
|
||||||
|
r.releaseForwarding()
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// releaseForwarding drops one IP forwarding reference, logging any error.
|
||||||
|
func (r *family) releaseForwarding() {
|
||||||
|
if err := r.ipFwdState.ReleaseForwarding(r.isV6()); err != nil {
|
||||||
|
log.Errorf("release IP forwarding: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// isV6 reports whether this family handles the IPv6 table.
|
||||||
|
func (r *family) isV6() bool {
|
||||||
|
return r.af.tableFamily == nftables.TableFamilyIPv6
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) AddInboundDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
||||||
|
ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
||||||
|
|
||||||
|
if _, exists := r.rules[ruleID]; exists {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
protoNum, err := r.af.protoNum(protocol)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("convert protocol to number: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs := []expr.Any{
|
||||||
|
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname(r.wgIface.Name()),
|
||||||
|
},
|
||||||
|
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 2},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 2,
|
||||||
|
Data: []byte{protoNum},
|
||||||
|
},
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 3,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 2,
|
||||||
|
Len: 2,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 3,
|
||||||
|
Data: binaryutil.BigEndian.PutUint16(originalPort),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
bits := 32
|
||||||
|
if localAddr.Is6() {
|
||||||
|
bits = 128
|
||||||
|
}
|
||||||
|
exprs = append(exprs, prefixMatchExprs(r.af, netip.PrefixFrom(localAddr, bits), false)...)
|
||||||
|
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 1,
|
||||||
|
Data: localAddr.AsSlice(),
|
||||||
|
},
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 2,
|
||||||
|
Data: binaryutil.BigEndian.PutUint16(translatedPort),
|
||||||
|
},
|
||||||
|
&expr.NAT{
|
||||||
|
Type: expr.NATTypeDestNAT,
|
||||||
|
Family: uint32(r.af.tableFamily),
|
||||||
|
RegAddrMin: 1,
|
||||||
|
RegProtoMin: 2,
|
||||||
|
RegProtoMax: 0,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
dnatRule := &nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chains[chainNameRoutingRdr],
|
||||||
|
Exprs: exprs,
|
||||||
|
UserData: []byte(ruleID),
|
||||||
|
}
|
||||||
|
r.conn.AddRule(dnatRule)
|
||||||
|
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf("add inbound DNAT rule: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.rules[ruleID] = dnatRule
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RemoveInboundDNAT removes an inbound DNAT rule.
|
||||||
|
func (r *family) RemoveInboundDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
||||||
|
if err := r.refreshRulesMap(); err != nil {
|
||||||
|
return fmt.Errorf(refreshRulesMapError, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
||||||
|
|
||||||
|
rule, exists := r.rules[ruleID]
|
||||||
|
if !exists {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if rule.Handle == 0 {
|
||||||
|
log.Warnf("inbound DNAT rule %s has no handle, removing stale entry", ruleID)
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.conn.DelRule(rule); err != nil {
|
||||||
|
return fmt.Errorf("delete inbound DNAT rule %s: %w", ruleID, err)
|
||||||
|
}
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf("flush delete inbound DNAT rule: %w", err)
|
||||||
|
}
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ensureNATOutputChain lazily creates the OUTPUT NAT chain on first use.
|
||||||
|
func (r *family) ensureNATOutputChain() error {
|
||||||
|
if _, exists := r.chains[chainNameNATOutput]; exists {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
r.chains[chainNameNATOutput] = r.conn.AddChain(&nftables.Chain{
|
||||||
|
Name: chainNameNATOutput,
|
||||||
|
Table: r.workTable,
|
||||||
|
Hooknum: nftables.ChainHookOutput,
|
||||||
|
Priority: nftables.ChainPriorityNATDest,
|
||||||
|
Type: nftables.ChainTypeNAT,
|
||||||
|
})
|
||||||
|
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
delete(r.chains, chainNameNATOutput)
|
||||||
|
return fmt.Errorf("create NAT output chain: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddOutputDNAT adds an OUTPUT chain DNAT rule for locally-generated traffic.
|
||||||
|
func (r *family) AddOutputDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
||||||
|
ruleID := firewall.RuleID(fmt.Sprintf("output-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
||||||
|
|
||||||
|
if _, exists := r.rules[ruleID]; exists {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.ensureNATOutputChain(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
protoNum, err := r.af.protoNum(protocol)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("convert protocol to number: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs := []expr.Any{
|
||||||
|
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: []byte{protoNum},
|
||||||
|
},
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 2,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 2,
|
||||||
|
Len: 2,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 2,
|
||||||
|
Data: binaryutil.BigEndian.PutUint16(originalPort),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
bits := 32
|
||||||
|
if localAddr.Is6() {
|
||||||
|
bits = 128
|
||||||
|
}
|
||||||
|
exprs = append(exprs, prefixMatchExprs(r.af, netip.PrefixFrom(localAddr, bits), false)...)
|
||||||
|
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 1,
|
||||||
|
Data: localAddr.AsSlice(),
|
||||||
|
},
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 2,
|
||||||
|
Data: binaryutil.BigEndian.PutUint16(translatedPort),
|
||||||
|
},
|
||||||
|
&expr.NAT{
|
||||||
|
Type: expr.NATTypeDestNAT,
|
||||||
|
Family: uint32(r.af.tableFamily),
|
||||||
|
RegAddrMin: 1,
|
||||||
|
RegProtoMin: 2,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
dnatRule := &nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chains[chainNameNATOutput],
|
||||||
|
Exprs: exprs,
|
||||||
|
UserData: []byte(ruleID),
|
||||||
|
}
|
||||||
|
r.conn.AddRule(dnatRule)
|
||||||
|
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf("add output DNAT rule: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.rules[ruleID] = dnatRule
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
|
||||||
|
func (r *family) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
||||||
|
if err := r.refreshRulesMap(); err != nil {
|
||||||
|
return fmt.Errorf(refreshRulesMapError, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ruleID := firewall.RuleID(fmt.Sprintf("output-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
||||||
|
|
||||||
|
rule, exists := r.rules[ruleID]
|
||||||
|
if !exists {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if rule.Handle == 0 {
|
||||||
|
log.Warnf("output DNAT rule %s has no handle, removing stale entry", ruleID)
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.conn.DelRule(rule); err != nil {
|
||||||
|
return fmt.Errorf("delete output DNAT rule %s: %w", ruleID, err)
|
||||||
|
}
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf("flush delete output DNAT rule: %w", err)
|
||||||
|
}
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -82,7 +82,7 @@ func dnatV6(port uint16) fw.ForwardRule {
|
|||||||
// v4 refcount at zero.
|
// v4 refcount at zero.
|
||||||
func TestNftablesDNAT_RefcountBalancedV4(t *testing.T) {
|
func TestNftablesDNAT_RefcountBalancedV4(t *testing.T) {
|
||||||
m := newNftRefcountManager(t, false)
|
m := newNftRefcountManager(t, false)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(dnatV4(8081))
|
r1, err := m.AddDNATRule(dnatV4(8081))
|
||||||
require.NoError(t, err, "add v4 dnat 1")
|
require.NoError(t, err, "add v4 dnat 1")
|
||||||
@@ -111,9 +111,9 @@ func TestNftablesDNAT_RefcountBalancedV4(t *testing.T) {
|
|||||||
// and decrements back to zero on Delete.
|
// and decrements back to zero on Delete.
|
||||||
func TestNftablesDNAT_RefcountBalancedV6(t *testing.T) {
|
func TestNftablesDNAT_RefcountBalancedV6(t *testing.T) {
|
||||||
m := newNftRefcountManager(t, true)
|
m := newNftRefcountManager(t, true)
|
||||||
require.NotNil(t, m.router6, "v6 router")
|
require.NotNil(t, m.family6, "v6 family")
|
||||||
require.Same(t, m.router.ipFwdState, m.router6.ipFwdState, "shared state")
|
require.Same(t, m.family4.ipFwdState, m.family6.ipFwdState, "shared state")
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(dnatV6(9091))
|
r1, err := m.AddDNATRule(dnatV6(9091))
|
||||||
require.NoError(t, err, "add v6 dnat 1")
|
require.NoError(t, err, "add v6 dnat 1")
|
||||||
@@ -142,7 +142,7 @@ func TestNftablesDNAT_RefcountBalancedV6(t *testing.T) {
|
|||||||
// ForwardRule) does not double-increment the refcount.
|
// ForwardRule) does not double-increment the refcount.
|
||||||
func TestNftablesDNAT_DuplicateAddNoLeak(t *testing.T) {
|
func TestNftablesDNAT_DuplicateAddNoLeak(t *testing.T) {
|
||||||
m := newNftRefcountManager(t, true)
|
m := newNftRefcountManager(t, true)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
rule := dnatV4(8083)
|
rule := dnatV4(8083)
|
||||||
r1, err := m.AddDNATRule(rule)
|
r1, err := m.AddDNATRule(rule)
|
||||||
@@ -165,7 +165,7 @@ func TestNftablesDNAT_DuplicateAddNoLeak(t *testing.T) {
|
|||||||
// never added does not underflow the refcount.
|
// never added does not underflow the refcount.
|
||||||
func TestNftablesDNAT_DeleteMissingNoUnderflow(t *testing.T) {
|
func TestNftablesDNAT_DeleteMissingNoUnderflow(t *testing.T) {
|
||||||
m := newNftRefcountManager(t, true)
|
m := newNftRefcountManager(t, true)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
// Construct a Rule reference for something never added. The router stores
|
// Construct a Rule reference for something never added. The router stores
|
||||||
// rules by ID(), and DeleteDNATRule looks them up in r.rules; a missing
|
// rules by ID(), and DeleteDNATRule looks them up in r.rules; a missing
|
||||||
@@ -195,7 +195,7 @@ func TestNftablesDNAT_DeleteMissingNoUnderflow(t *testing.T) {
|
|||||||
// and a single DisableRouting drops both back to zero.
|
// and a single DisableRouting drops both back to zero.
|
||||||
func TestNftablesRouting_RepeatedEnableSingleReference(t *testing.T) {
|
func TestNftablesRouting_RepeatedEnableSingleReference(t *testing.T) {
|
||||||
m := newNftRefcountManager(t, true)
|
m := newNftRefcountManager(t, true)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
require.NoError(t, m.EnableRouting(), "first enable")
|
require.NoError(t, m.EnableRouting(), "first enable")
|
||||||
require.NoError(t, m.EnableRouting(), "second enable")
|
require.NoError(t, m.EnableRouting(), "second enable")
|
||||||
@@ -214,7 +214,7 @@ func TestNftablesRouting_RepeatedEnableSingleReference(t *testing.T) {
|
|||||||
// DisableRouting does not release references held by active DNAT rules.
|
// DisableRouting does not release references held by active DNAT rules.
|
||||||
func TestNftablesRouting_DisableKeepsDNATReference(t *testing.T) {
|
func TestNftablesRouting_DisableKeepsDNATReference(t *testing.T) {
|
||||||
m := newNftRefcountManager(t, true)
|
m := newNftRefcountManager(t, true)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(dnatV6(9095))
|
r1, err := m.AddDNATRule(dnatV6(9095))
|
||||||
require.NoError(t, err, "add v6 dnat")
|
require.NoError(t, err, "add v6 dnat")
|
||||||
@@ -232,7 +232,7 @@ func TestNftablesRouting_DisableKeepsDNATReference(t *testing.T) {
|
|||||||
// twice does not underflow the refcount (the second delete is a no-op).
|
// twice does not underflow the refcount (the second delete is a no-op).
|
||||||
func TestNftablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) {
|
func TestNftablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) {
|
||||||
m := newNftRefcountManager(t, true)
|
m := newNftRefcountManager(t, true)
|
||||||
state := m.router.ipFwdState
|
state := m.family4.ipFwdState
|
||||||
|
|
||||||
r1, err := m.AddDNATRule(dnatV6(9093))
|
r1, err := m.AddDNATRule(dnatV6(9093))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|||||||
@@ -0,0 +1,249 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/coreos/go-iptables/iptables"
|
||||||
|
"github.com/google/nftables"
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
"github.com/netbirdio/netbird/client/firewall/firewalld"
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/routemanager/ipfwdstate"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/routemanager/refcounter"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
tableNat = "nat"
|
||||||
|
tableMangle = "mangle"
|
||||||
|
tableRaw = "raw"
|
||||||
|
tableSecurity = "security"
|
||||||
|
|
||||||
|
chainNameNatPrerouting = "PREROUTING"
|
||||||
|
chainNameRoutingFw = "netbird-rt-fwd"
|
||||||
|
chainNameRoutingNat = "netbird-rt-postrouting"
|
||||||
|
chainNameRoutingRdr = "netbird-rt-redirect"
|
||||||
|
chainNameNATOutput = "netbird-nat-output"
|
||||||
|
chainNameForward = "FORWARD"
|
||||||
|
chainNameMangleForward = "netbird-mangle-forward"
|
||||||
|
|
||||||
|
// Peer ACL chain names.
|
||||||
|
chainNameInputRules = "netbird-acl-input-rules"
|
||||||
|
chainNameInputFilter = "netbird-acl-input-filter"
|
||||||
|
chainNameForwardFilter = "netbird-acl-forward-filter"
|
||||||
|
chainNameManglePrerouting = "netbird-mangle-prerouting"
|
||||||
|
chainNameManglePostrouting = "netbird-mangle-postrouting"
|
||||||
|
|
||||||
|
flushError = "flush: %w"
|
||||||
|
|
||||||
|
firewalldTableName = "firewalld"
|
||||||
|
|
||||||
|
userDataAcceptForwardRuleIif = "frwacceptiif"
|
||||||
|
userDataAcceptForwardRuleOif = "frwacceptoif"
|
||||||
|
userDataAcceptInputRule = "inputaccept"
|
||||||
|
|
||||||
|
dnatSuffix firewall.RuleID = "_dnat"
|
||||||
|
snatSuffix firewall.RuleID = "_snat"
|
||||||
|
|
||||||
|
// ipv4TCPHeaderSize is the minimum IPv4 (20) + TCP (20) header size for MSS calculation.
|
||||||
|
ipv4TCPHeaderSize = 40
|
||||||
|
// ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation.
|
||||||
|
ipv6TCPHeaderSize = 60
|
||||||
|
|
||||||
|
// maxPrefixesSet 1638 prefixes start to fail, taking some margin
|
||||||
|
maxPrefixesSet = 1500
|
||||||
|
refreshRulesMapError = "refresh rules map: %w"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
errFilterTableNotFound = fmt.Errorf("'filter' table not found")
|
||||||
|
)
|
||||||
|
|
||||||
|
type setInput struct {
|
||||||
|
set firewall.Set
|
||||||
|
prefixes []netip.Prefix
|
||||||
|
}
|
||||||
|
|
||||||
|
// family holds the per-address-family nftables state. One instance
|
||||||
|
// handles route ACLs, peer ACLs, NAT, DNAT, and MSS clamping for a
|
||||||
|
// single family; the top-level Manager owns one for v4 and another
|
||||||
|
// for v6. The name predates the peer-ACL absorption; it's effectively
|
||||||
|
// the per-family backend now.
|
||||||
|
type family struct {
|
||||||
|
conn *nftables.Conn
|
||||||
|
workTable *nftables.Table
|
||||||
|
filterTable *nftables.Table
|
||||||
|
chains map[string]*nftables.Chain
|
||||||
|
|
||||||
|
// filters holds peer + route filter rules keyed by content hash.
|
||||||
|
// AddFilterRule writes here; DeleteFilterRule looks up by id.
|
||||||
|
filters map[firewall.RuleID]*Rule
|
||||||
|
|
||||||
|
// rules holds NAT, DNAT, and external accept rules (auxiliary
|
||||||
|
// plumbing that isn't a filter rule).
|
||||||
|
rules map[firewall.RuleID]*nftables.Rule
|
||||||
|
|
||||||
|
// Peer ACL chain handles.
|
||||||
|
chainInputRules *nftables.Chain
|
||||||
|
chainPrerouting *nftables.Chain
|
||||||
|
routingFwChainName string
|
||||||
|
|
||||||
|
ipsetCounter *refcounter.Counter[string, setInput, *nftables.Set]
|
||||||
|
|
||||||
|
af addrFamily
|
||||||
|
wgIface iFaceMapper
|
||||||
|
ipFwdState *ipfwdstate.IPForwardingState
|
||||||
|
legacyManagement bool
|
||||||
|
mtu uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFamily(workTable *nftables.Table, wgIface iFaceMapper, mtu uint16) *family {
|
||||||
|
r := &family{
|
||||||
|
conn: &nftables.Conn{},
|
||||||
|
workTable: workTable,
|
||||||
|
chains: make(map[string]*nftables.Chain),
|
||||||
|
filters: make(map[firewall.RuleID]*Rule),
|
||||||
|
rules: make(map[firewall.RuleID]*nftables.Rule),
|
||||||
|
routingFwChainName: chainNameRoutingFw,
|
||||||
|
af: familyForAddr(workTable.Family == nftables.TableFamilyIPv4),
|
||||||
|
wgIface: wgIface,
|
||||||
|
ipFwdState: ipfwdstate.NewIPForwardingState(wgIface.Name()),
|
||||||
|
mtu: mtu,
|
||||||
|
}
|
||||||
|
|
||||||
|
r.ipsetCounter = refcounter.New(
|
||||||
|
r.createIpSet,
|
||||||
|
r.deleteIpSet,
|
||||||
|
)
|
||||||
|
|
||||||
|
var err error
|
||||||
|
r.filterTable, err = r.loadFilterTable()
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("ip filter table not found: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) init(workTable *nftables.Table) error {
|
||||||
|
r.workTable = workTable
|
||||||
|
|
||||||
|
if err := r.removeAcceptFilterRules(); err != nil {
|
||||||
|
log.Errorf("failed to clean up rules from filter table: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.createContainers(); err != nil {
|
||||||
|
return fmt.Errorf("create containers: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.setupDataPlaneMark(); err != nil {
|
||||||
|
log.Errorf("failed to set up data plane mark: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.createDefaultChains(); err != nil {
|
||||||
|
return fmt.Errorf("create default acl chains: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reset cleans existing nftables filter table rules from the system
|
||||||
|
func (r *family) Reset() error {
|
||||||
|
// clear without deleting the ipsets, the nf table will be deleted by the caller
|
||||||
|
r.ipsetCounter.Clear()
|
||||||
|
|
||||||
|
var merr *multierror.Error
|
||||||
|
|
||||||
|
if err := r.removeAcceptFilterRules(); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove accept filter rules: %w", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := firewalld.UntrustInterface(r.wgIface.Name()); err != nil {
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) loadFilterTable() (*nftables.Table, error) {
|
||||||
|
tables, err := r.conn.ListTablesOfFamily(r.af.tableFamily)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("list tables: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, table := range tables {
|
||||||
|
if table.Name == "filter" {
|
||||||
|
return table, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, errFilterTableNotFound
|
||||||
|
}
|
||||||
|
|
||||||
|
func hookName(hook *nftables.ChainHook) string {
|
||||||
|
if hook == nil {
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
switch *hook {
|
||||||
|
case *nftables.ChainHookForward:
|
||||||
|
return chainNameForward
|
||||||
|
case *nftables.ChainHookInput:
|
||||||
|
return chainNameInput
|
||||||
|
default:
|
||||||
|
return fmt.Sprintf("hook(%d)", *hook)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func familyName(family nftables.TableFamily) string {
|
||||||
|
switch family {
|
||||||
|
case nftables.TableFamilyIPv4:
|
||||||
|
return "ip"
|
||||||
|
case nftables.TableFamilyIPv6:
|
||||||
|
return "ip6"
|
||||||
|
case nftables.TableFamilyINet:
|
||||||
|
return "inet"
|
||||||
|
default:
|
||||||
|
return fmt.Sprintf("family(%d)", family)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) iptablesProto() iptables.Protocol {
|
||||||
|
if r.af.tableFamily == nftables.TableFamilyIPv6 {
|
||||||
|
return iptables.ProtocolIPv6
|
||||||
|
}
|
||||||
|
return iptables.ProtocolIPv4
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) refreshRulesMap() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
newRules := make(map[firewall.RuleID]*nftables.Rule)
|
||||||
|
for _, chain := range r.chains {
|
||||||
|
rules, err := r.conn.GetRules(chain.Table, chain)
|
||||||
|
if err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("list rules for chain %s: %w", chain.Name, err))
|
||||||
|
// preserve existing entries for this chain since we can't verify their state
|
||||||
|
for k, v := range r.rules {
|
||||||
|
if v.Chain != nil && v.Chain.Name == chain.Name {
|
||||||
|
newRules[k] = v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, rule := range rules {
|
||||||
|
if len(rule.UserData) > 0 {
|
||||||
|
newRules[firewall.RuleID(rule.UserData)] = rule
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
r.rules = newRules
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
@@ -0,0 +1,540 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
|
||||||
|
"github.com/google/nftables"
|
||||||
|
"github.com/google/nftables/binaryutil"
|
||||||
|
"github.com/google/nftables/expr"
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
nbid "github.com/netbirdio/netbird/client/internal/acl/id"
|
||||||
|
)
|
||||||
|
|
||||||
|
// AddFilterRule installs one nftables packet-filter rule. With
|
||||||
|
// destination empty the rule goes to the peer ACL input chain plus a
|
||||||
|
// paired prerouting mangle rule for the redirect mark. With
|
||||||
|
// destination set (prefix or named set) it goes to the route ACL
|
||||||
|
// forward chain. Multi-source rules collapse to one nftables rule
|
||||||
|
// backed by the shared refcounted hash:net set.
|
||||||
|
func (r *family) AddFilterRule(
|
||||||
|
id []byte,
|
||||||
|
sources []netip.Prefix,
|
||||||
|
destination firewall.Network,
|
||||||
|
proto firewall.Protocol,
|
||||||
|
sPort *firewall.Port,
|
||||||
|
dPort *firewall.Port,
|
||||||
|
action firewall.Action,
|
||||||
|
) (firewall.Rule, error) {
|
||||||
|
isRoute := !destination.IsZero()
|
||||||
|
|
||||||
|
ruleID := nbid.GenerateRuleID(sources, destination, proto, sPort, dPort, action)
|
||||||
|
if existing, ok := r.filters[ruleID]; ok {
|
||||||
|
return existing, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
srcExprs, err := r.applyNetwork(sourceNetwork(sources), sources, true)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("apply source: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var exprs []expr.Any
|
||||||
|
if isRoute {
|
||||||
|
exprs, err = r.buildRouteFilterExprs(srcExprs, destination, proto, sPort, dPort)
|
||||||
|
} else {
|
||||||
|
exprs, err = r.buildPeerFilterExprs(srcExprs, proto, sPort, dPort)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
r.dropNetworkMatch(srcExprs)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
mainExprs := slices.Clone(exprs)
|
||||||
|
verdict := expr.VerdictAccept
|
||||||
|
if action == firewall.ActionDrop {
|
||||||
|
verdict = expr.VerdictDrop
|
||||||
|
}
|
||||||
|
mainExprs = append(mainExprs, &expr.Verdict{Kind: verdict})
|
||||||
|
|
||||||
|
chain := r.chainInputRules
|
||||||
|
if isRoute {
|
||||||
|
chain = r.chains[chainNameRoutingFw]
|
||||||
|
}
|
||||||
|
|
||||||
|
userData := []byte(ruleID)
|
||||||
|
|
||||||
|
// Build the paired prerouting mangle rule before flushing so both
|
||||||
|
// rules commit in one transaction. An anonymous port set binds to
|
||||||
|
// exactly one rule, so the mangle rule needs its own expression list
|
||||||
|
// with fresh sets, not a clone of the main rule's. Guard on the
|
||||||
|
// prerouting chain first: building the expressions queues the port
|
||||||
|
// set, so skipping the build when there is no chain to bind it to
|
||||||
|
// keeps an unbound set out of the connection batch.
|
||||||
|
var mangleRule *nftables.Rule
|
||||||
|
if !isRoute && r.chainPrerouting != nil {
|
||||||
|
mangleExprs, err := r.buildPeerFilterExprs(srcExprs, proto, sPort, dPort)
|
||||||
|
if err != nil {
|
||||||
|
r.dropNetworkMatch(exprs)
|
||||||
|
return nil, fmt.Errorf("build mangle rule: %w", err)
|
||||||
|
}
|
||||||
|
mangleRule = r.queuePreroutingRule(mangleExprs, userData)
|
||||||
|
}
|
||||||
|
|
||||||
|
nftRule := &nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: chain,
|
||||||
|
Exprs: mainExprs,
|
||||||
|
UserData: userData,
|
||||||
|
}
|
||||||
|
if action == firewall.ActionDrop {
|
||||||
|
nftRule = r.conn.InsertRule(nftRule)
|
||||||
|
} else {
|
||||||
|
nftRule = r.conn.AddRule(nftRule)
|
||||||
|
}
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
r.dropNetworkMatch(exprs)
|
||||||
|
return nil, fmt.Errorf(flushError, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rule := &Rule{
|
||||||
|
nftRule: nftRule,
|
||||||
|
mangleRule: mangleRule,
|
||||||
|
sources: sources,
|
||||||
|
id: ruleID,
|
||||||
|
}
|
||||||
|
r.filters[ruleID] = rule
|
||||||
|
|
||||||
|
log.Debugf("added filter rule: sources=%v, destination=%v, proto=%v, sPort=%v, dPort=%v, action=%v",
|
||||||
|
sources, destination, proto, sPort, dPort, action)
|
||||||
|
return rule, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildPeerFilterExprs assembles the input-chain (peer ACL) match: the
|
||||||
|
// IP-header protocol byte read via Payload, then source, then ports
|
||||||
|
// (no counter), matching the historical peer shape so per-rule kernel
|
||||||
|
// state is identical to pre-unification.
|
||||||
|
func (r *family) buildPeerFilterExprs(
|
||||||
|
srcExprs []expr.Any,
|
||||||
|
proto firewall.Protocol,
|
||||||
|
sPort, dPort *firewall.Port,
|
||||||
|
) ([]expr.Any, error) {
|
||||||
|
var exprs []expr.Any
|
||||||
|
|
||||||
|
if proto != firewall.ProtocolALL {
|
||||||
|
protoNum, err := r.af.protoNum(proto)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("convert protocol to number: %w", err)
|
||||||
|
}
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: r.af.protoOffset,
|
||||||
|
Len: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{protoNum}},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
exprs = append(exprs, srcExprs...)
|
||||||
|
|
||||||
|
portExprs, err := r.applyPorts(sPort, dPort)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
exprs = append(exprs, portExprs...)
|
||||||
|
return exprs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildRouteFilterExprs assembles the forward-chain (route ACL) match:
|
||||||
|
// source, then destination, then optional proto/ports, then a counter.
|
||||||
|
func (r *family) buildRouteFilterExprs(
|
||||||
|
srcExprs []expr.Any,
|
||||||
|
destination firewall.Network,
|
||||||
|
proto firewall.Protocol,
|
||||||
|
sPort, dPort *firewall.Port,
|
||||||
|
) ([]expr.Any, error) {
|
||||||
|
exprs := append([]expr.Any{}, srcExprs...)
|
||||||
|
|
||||||
|
destExprs, err := r.applyNetwork(destination, nil, false)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("apply destination: %w", err)
|
||||||
|
}
|
||||||
|
exprs = append(exprs, destExprs...)
|
||||||
|
|
||||||
|
if proto != firewall.ProtocolALL {
|
||||||
|
protoNum, err := r.af.protoNum(proto)
|
||||||
|
if err != nil {
|
||||||
|
r.dropNetworkMatch(destExprs)
|
||||||
|
return nil, fmt.Errorf("convert protocol to number: %w", err)
|
||||||
|
}
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{protoNum}},
|
||||||
|
)
|
||||||
|
|
||||||
|
portExprs, err := r.applyPorts(sPort, dPort)
|
||||||
|
if err != nil {
|
||||||
|
r.dropNetworkMatch(destExprs)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
exprs = append(exprs, portExprs...)
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs = append(exprs, &expr.Counter{})
|
||||||
|
return exprs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) hasRule(id firewall.RuleID) bool {
|
||||||
|
_, ok := r.filters[id]
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) hasDNATRule(id firewall.RuleID) bool {
|
||||||
|
_, ok := r.rules[id+dnatSuffix]
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteFilterRule removes a previously installed filter rule. Source
|
||||||
|
// set references are recovered from the stored rule's expressions via
|
||||||
|
// findSets and dropped from the shared refcounter.
|
||||||
|
func (r *family) DeleteFilterRule(rule firewall.Rule) error {
|
||||||
|
ruleID := rule.ID()
|
||||||
|
pr, ok := r.filters[ruleID]
|
||||||
|
if !ok {
|
||||||
|
log.Debugf("filter rule %s not found", ruleID)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// A freshly added rule carries no handle until it is read back from
|
||||||
|
// the kernel, and Flush only refreshes the peer chains. Pull live
|
||||||
|
// handles for this rule's chain before deciding it is stale so route
|
||||||
|
// rules (which Flush never refreshes) can actually be deleted. A
|
||||||
|
// refresh failure aborts the delete without touching tracking state,
|
||||||
|
// so the caller can retry while the rule may still exist in the kernel.
|
||||||
|
if pr.nftRule.Handle == 0 {
|
||||||
|
if err := r.refreshRuleHandles(pr.nftRule.Chain, false); err != nil {
|
||||||
|
return fmt.Errorf("refresh handles for chain %s: %w", pr.nftRule.Chain.Name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Refresh the mangle handle independently: the main rule's handle can
|
||||||
|
// be populated while the prerouting refresh during Flush failed, and
|
||||||
|
// gating the mangle refresh on the main handle would leak the mangle
|
||||||
|
// rule on delete.
|
||||||
|
if pr.mangleRule != nil && pr.mangleRule.Handle == 0 {
|
||||||
|
if err := r.refreshRuleHandles(r.chainPrerouting, true); err != nil {
|
||||||
|
return fmt.Errorf("refresh mangle handles: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if pr.nftRule.Handle == 0 {
|
||||||
|
log.Warnf("filter rule %s has no handle, removing stale entry", ruleID)
|
||||||
|
// The paired mangle rule can still be in the kernel with a live
|
||||||
|
// handle. Dropping the tracking entry without removing it would
|
||||||
|
// leave a prerouting rule that nothing can find again.
|
||||||
|
if err := r.deleteMangleRule(pr, ruleID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
r.dropNetworkMatch(pr.nftRule.Exprs)
|
||||||
|
delete(r.filters, ruleID)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.conn.DelRule(pr.nftRule); err != nil {
|
||||||
|
log.Errorf("queue rule delete: %v", err)
|
||||||
|
}
|
||||||
|
r.queueMangleDelete(pr)
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf("flush delete %s: %w", ruleID, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.dropNetworkMatch(pr.nftRule.Exprs)
|
||||||
|
delete(r.filters, ruleID)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// deleteMangleRule removes the prerouting rule paired with a filter rule on
|
||||||
|
// its own, for the paths that drop the filter rule's tracking without queueing
|
||||||
|
// a delete for it.
|
||||||
|
func (r *family) deleteMangleRule(pr *Rule, ruleID firewall.RuleID) error {
|
||||||
|
if pr.mangleRule == nil || pr.mangleRule.Handle == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
r.queueMangleDelete(pr)
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf("flush mangle delete %s: %w", ruleID, err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// queueMangleDelete queues the delete of the rule's prerouting counterpart, if
|
||||||
|
// it has one. The caller commits it.
|
||||||
|
func (r *family) queueMangleDelete(pr *Rule) {
|
||||||
|
if pr.mangleRule == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := r.conn.DelRule(pr.mangleRule); err != nil {
|
||||||
|
log.Errorf("queue mangle rule delete: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) decrementSetCounter(rule *nftables.Rule) error {
|
||||||
|
if r.ipsetCounter == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
sets := findSets(rule)
|
||||||
|
|
||||||
|
var merr *multierror.Error
|
||||||
|
for _, setName := range sets {
|
||||||
|
if _, err := r.ipsetCounter.Decrement(setName); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("decrement set counter: %w", err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// dropNetworkMatch undoes whatever the source/destination match
|
||||||
|
// reserved. Safe to call when the spec is empty or holds only inline
|
||||||
|
// matchers.
|
||||||
|
func (r *family) dropNetworkMatch(exprs []expr.Any) {
|
||||||
|
if r.ipsetCounter == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
for _, e := range exprs {
|
||||||
|
lookup, ok := e.(*expr.Lookup)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, err := r.ipsetCounter.Decrement(lookup.SetName); err != nil {
|
||||||
|
log.Errorf("rollback ipset decrement %s: %v", lookup.SetName, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) applyNetwork(
|
||||||
|
network firewall.Network,
|
||||||
|
setPrefixes []netip.Prefix,
|
||||||
|
isSource bool,
|
||||||
|
) ([]expr.Any, error) {
|
||||||
|
if network.IsSet() {
|
||||||
|
exprs, err := r.getIpSet(network.Set, setPrefixes, isSource)
|
||||||
|
if err != nil {
|
||||||
|
side := "destination"
|
||||||
|
if isSource {
|
||||||
|
side = "source"
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("%s set: %w", side, err)
|
||||||
|
}
|
||||||
|
return exprs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if network.IsPrefix() {
|
||||||
|
return prefixMatchExprs(r.af, network.Prefix, isSource), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// applyPort builds the transport-header port match. A single value
|
||||||
|
// compares directly, a range uses a range expression, and multiple
|
||||||
|
// values go through an anonymous constant set: consecutive cmp
|
||||||
|
// expressions AND together, so chained equality comparisons could
|
||||||
|
// never match more than one port. The set is queued on the
|
||||||
|
// connection and committed by the caller's flush together with the
|
||||||
|
// rule that binds it.
|
||||||
|
func (r *family) applyPort(port *firewall.Port, isSource bool) ([]expr.Any, error) {
|
||||||
|
if port == nil || len(port.Values) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// dst port
|
||||||
|
offset := uint32(2)
|
||||||
|
if isSource {
|
||||||
|
// src port
|
||||||
|
offset = 0
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs := []expr.Any{
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: offset,
|
||||||
|
Len: 2,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case port.IsRange && len(port.Values) == 2:
|
||||||
|
exprs = append(exprs, &expr.Range{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
FromData: binaryutil.BigEndian.PutUint16(port.Values[0]),
|
||||||
|
ToData: binaryutil.BigEndian.PutUint16(port.Values[1]),
|
||||||
|
})
|
||||||
|
case len(port.Values) == 1:
|
||||||
|
exprs = append(exprs, &expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: binaryutil.BigEndian.PutUint16(port.Values[0]),
|
||||||
|
})
|
||||||
|
default:
|
||||||
|
lookup, err := r.anonymousPortSet(port.Values)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
exprs = append(exprs, lookup)
|
||||||
|
}
|
||||||
|
|
||||||
|
return exprs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// anonymousPortSet queues an anonymous constant set holding the given
|
||||||
|
// ports on the connection and returns a lookup against it. The set is
|
||||||
|
// committed by the caller's flush together with the rule that binds it.
|
||||||
|
func (r *family) anonymousPortSet(values []uint16) (*expr.Lookup, error) {
|
||||||
|
set := &nftables.Set{
|
||||||
|
Anonymous: true,
|
||||||
|
Constant: true,
|
||||||
|
Table: r.workTable,
|
||||||
|
KeyType: nftables.TypeInetService,
|
||||||
|
}
|
||||||
|
elements := make([]nftables.SetElement, 0, len(values))
|
||||||
|
for _, p := range values {
|
||||||
|
elements = append(elements, nftables.SetElement{Key: binaryutil.BigEndian.PutUint16(p)})
|
||||||
|
}
|
||||||
|
if err := r.conn.AddSet(set, elements); err != nil {
|
||||||
|
return nil, fmt.Errorf("add anonymous port set: %w", err)
|
||||||
|
}
|
||||||
|
return &expr.Lookup{
|
||||||
|
SourceRegister: 1,
|
||||||
|
SetID: set.ID,
|
||||||
|
SetName: set.Name,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// applyPorts builds the source then destination port matches.
|
||||||
|
func (r *family) applyPorts(sPort, dPort *firewall.Port) ([]expr.Any, error) {
|
||||||
|
sPortExprs, err := r.applyPort(sPort, true)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("apply source port: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
dPortExprs, err := r.applyPort(dPort, false)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("apply destination port: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return append(sPortExprs, dPortExprs...), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// prefixMatchExprs is the family-aware match sequence for a CIDR
|
||||||
|
// prefix. /0 returns nil; a host prefix (full bit length for the
|
||||||
|
// family) skips the bitwise step since the mask is all-ones. Shared
|
||||||
|
// between family and aclManager so both treat single prefixes
|
||||||
|
// identically.
|
||||||
|
func prefixMatchExprs(af addrFamily, prefix netip.Prefix, isSource bool) []expr.Any {
|
||||||
|
offset := af.dstAddrOffset
|
||||||
|
if isSource {
|
||||||
|
offset = af.srcAddrOffset
|
||||||
|
}
|
||||||
|
|
||||||
|
ones := prefix.Bits()
|
||||||
|
if ones == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
payload := &expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: offset,
|
||||||
|
Len: af.addrLen,
|
||||||
|
}
|
||||||
|
cmp := &expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: prefix.Masked().Addr().AsSlice(),
|
||||||
|
}
|
||||||
|
|
||||||
|
if ones == af.totalBits {
|
||||||
|
return []expr.Any{payload, cmp}
|
||||||
|
}
|
||||||
|
|
||||||
|
mask := net.CIDRMask(ones, af.totalBits)
|
||||||
|
xor := make([]byte, af.addrLen)
|
||||||
|
return []expr.Any{
|
||||||
|
payload,
|
||||||
|
&expr.Bitwise{
|
||||||
|
DestRegister: 1,
|
||||||
|
SourceRegister: 1,
|
||||||
|
Len: af.addrLen,
|
||||||
|
Mask: mask,
|
||||||
|
Xor: xor,
|
||||||
|
},
|
||||||
|
cmp,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func getCtNewExprs() []expr.Any {
|
||||||
|
return []expr.Any{
|
||||||
|
&expr.Ct{
|
||||||
|
Key: expr.CtKeySTATE,
|
||||||
|
Register: 1,
|
||||||
|
},
|
||||||
|
&expr.Bitwise{
|
||||||
|
SourceRegister: 1,
|
||||||
|
DestRegister: 1,
|
||||||
|
Len: 4,
|
||||||
|
Mask: binaryutil.NativeEndian.PutUint32(expr.CtStateBitNEW),
|
||||||
|
Xor: binaryutil.NativeEndian.PutUint32(0),
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpNeq,
|
||||||
|
Register: 1,
|
||||||
|
Data: []byte{0, 0, 0, 0},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// sourceNetwork classifies a source-prefix list into the firewall.Network
|
||||||
|
// shape the rest of the spec-builder consumes: empty for match-any, a
|
||||||
|
// single prefix inline, or an ipset for multiple sources.
|
||||||
|
func sourceNetwork(sources []netip.Prefix) firewall.Network {
|
||||||
|
switch {
|
||||||
|
case len(sources) == 0:
|
||||||
|
return firewall.Network{}
|
||||||
|
case len(sources) == 1 && sources[0].Bits() == 0:
|
||||||
|
return firewall.Network{}
|
||||||
|
case len(sources) == 1:
|
||||||
|
return firewall.Network{Prefix: sources[0]}
|
||||||
|
default:
|
||||||
|
return firewall.Network{Set: firewall.NewPrefixSet(sources)}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func ifname(n string) []byte {
|
||||||
|
b := make([]byte, 16)
|
||||||
|
copy(b, n+"\x00")
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
// findSets scans an nftables rule's expressions for expr.Lookup and
|
||||||
|
// returns the named sets in occurrence order. Used at delete time to
|
||||||
|
// drop ipsetCounter references; peer and route ACLs go through it.
|
||||||
|
func findSets(rule *nftables.Rule) []string {
|
||||||
|
var sets []string
|
||||||
|
for _, e := range rule.Exprs {
|
||||||
|
if lookup, ok := e.(*expr.Lookup); ok {
|
||||||
|
sets = append(sets, lookup.SetName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return sets
|
||||||
|
}
|
||||||
@@ -0,0 +1,90 @@
|
|||||||
|
//go:build privileged
|
||||||
|
|
||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/google/nftables"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/iface"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestInterfaceAllowerInputOnly verifies the userspace-mode allower opens the
|
||||||
|
// interface on the INPUT hook of foreign chains only (not FORWARD, since the
|
||||||
|
// userspace router never forwards in the kernel), creates no netbird work
|
||||||
|
// table, and removes its rules on Close.
|
||||||
|
func TestInterfaceAllowerInputOnly(t *testing.T) {
|
||||||
|
if os.Geteuid() != 0 {
|
||||||
|
t.Skip("root required")
|
||||||
|
}
|
||||||
|
|
||||||
|
require.False(t, ipTableExists(t, getTableName()), "precondition: no stale netbird table")
|
||||||
|
|
||||||
|
conn := &nftables.Conn{}
|
||||||
|
extTable := conn.AddTable(&nftables.Table{Name: "nbtest_extchains", Family: nftables.TableFamilyINet})
|
||||||
|
inputChain := conn.AddChain(&nftables.Chain{
|
||||||
|
Name: "ext_input", Table: extTable,
|
||||||
|
Hooknum: nftables.ChainHookInput, Priority: nftables.ChainPriorityFilter, Type: nftables.ChainTypeFilter,
|
||||||
|
})
|
||||||
|
forwardChain := conn.AddChain(&nftables.Chain{
|
||||||
|
Name: "ext_forward", Table: extTable,
|
||||||
|
Hooknum: nftables.ChainHookForward, Priority: nftables.ChainPriorityFilter, Type: nftables.ChainTypeFilter,
|
||||||
|
})
|
||||||
|
require.NoError(t, conn.Flush(), "create external table and chains")
|
||||||
|
t.Cleanup(func() {
|
||||||
|
c := &nftables.Conn{}
|
||||||
|
c.DelTable(extTable)
|
||||||
|
_ = c.Flush()
|
||||||
|
})
|
||||||
|
|
||||||
|
allower, err := NewInterfaceAllower(ifaceMock, iface.DefaultMTU)
|
||||||
|
require.NoError(t, err, "create allower")
|
||||||
|
require.NoError(t, allower.Apply(), "apply")
|
||||||
|
|
||||||
|
require.True(t, chainHasUserData(t, extTable, inputChain, userDataAcceptInputRule),
|
||||||
|
"external INPUT chain should get the accept rule")
|
||||||
|
require.Len(t, listRules(t, extTable, forwardChain), 0,
|
||||||
|
"external FORWARD chain must not be opened in userspace mode")
|
||||||
|
require.False(t, ipTableExists(t, getTableName()),
|
||||||
|
"allower must not create a netbird work table")
|
||||||
|
|
||||||
|
require.NoError(t, allower.Close(), "close")
|
||||||
|
require.False(t, chainHasUserData(t, extTable, inputChain, userDataAcceptInputRule),
|
||||||
|
"accept rule should be removed on close")
|
||||||
|
}
|
||||||
|
|
||||||
|
func listRules(t *testing.T, table *nftables.Table, chain *nftables.Chain) []*nftables.Rule {
|
||||||
|
t.Helper()
|
||||||
|
c := &nftables.Conn{}
|
||||||
|
rules, err := c.GetRules(table, chain)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return rules
|
||||||
|
}
|
||||||
|
|
||||||
|
func chainHasUserData(t *testing.T, table *nftables.Table, chain *nftables.Chain, ud string) bool {
|
||||||
|
for _, r := range listRules(t, table, chain) {
|
||||||
|
if bytes.Equal(r.UserData, []byte(ud)) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func ipTableExists(t *testing.T, name string) bool {
|
||||||
|
t.Helper()
|
||||||
|
c := &nftables.Conn{}
|
||||||
|
for _, fam := range []nftables.TableFamily{nftables.TableFamilyIPv4, nftables.TableFamilyIPv6} {
|
||||||
|
tbls, err := c.ListTablesOfFamily(fam)
|
||||||
|
require.NoError(t, err)
|
||||||
|
for _, tb := range tbls {
|
||||||
|
if tb.Name == name {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
@@ -0,0 +1,107 @@
|
|||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/google/nftables"
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
)
|
||||||
|
|
||||||
|
// InterfaceAllower opens the NetBird interface in the kernel's filter table and
|
||||||
|
// external chains and keeps them reconciled via a netlink monitor, so the host
|
||||||
|
// firewall doesn't drop traffic the NetBird firewall handles. It is used by the
|
||||||
|
// userspace firewall, where routing happens in the forwarder, so only INPUT is
|
||||||
|
// opened (the userspace router never forwards in the kernel).
|
||||||
|
//
|
||||||
|
// It owns its own families/connection and never creates a netbird work table.
|
||||||
|
// firewalld trust is handled by the caller, not here. Its operations are serial
|
||||||
|
// (Apply before the monitor starts; reconciles run on the single monitor
|
||||||
|
// goroutine; Close stops the monitor before removing), so it needs no locking.
|
||||||
|
//
|
||||||
|
// TODO: this opens nftables and the iptables-nft filter table (detected via
|
||||||
|
// nft), but not a legacy-iptables ruleset running in parallel with nftables.
|
||||||
|
// Such a host would keep its legacy filter chains closed for the interface.
|
||||||
|
type InterfaceAllower struct {
|
||||||
|
family4 *family
|
||||||
|
family6 *family
|
||||||
|
extMonitor *externalChainMonitor
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewInterfaceAllower builds an allower for the given interface. It returns an
|
||||||
|
// error when nftables is unavailable (e.g. an iptables-legacy host), so the
|
||||||
|
// caller can fall back to firewalld trust.
|
||||||
|
func NewInterfaceAllower(wgIface iFaceMapper, mtu uint16) (*InterfaceAllower, error) {
|
||||||
|
tableName := getTableName()
|
||||||
|
|
||||||
|
family4 := newFamily(&nftables.Table{Name: tableName, Family: nftables.TableFamilyIPv4}, wgIface, mtu)
|
||||||
|
|
||||||
|
// Probe nftables availability before committing to this backend.
|
||||||
|
if _, err := family4.conn.ListChainsOfTableFamily(nftables.TableFamilyINet); err != nil {
|
||||||
|
return nil, fmt.Errorf("nftables not available: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
a := &InterfaceAllower{family4: family4}
|
||||||
|
|
||||||
|
if wgIface.Address().HasIPv6() {
|
||||||
|
a.family6 = newFamily(&nftables.Table{Name: tableName, Family: nftables.TableFamilyIPv6}, wgIface, mtu)
|
||||||
|
}
|
||||||
|
|
||||||
|
a.extMonitor = newExternalChainMonitor(a)
|
||||||
|
return a, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply opens the interface (INPUT only) in the foreign filter chains and starts
|
||||||
|
// reconciling them on nftables changes.
|
||||||
|
func (a *InterfaceAllower) Apply() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
for _, f := range a.families() {
|
||||||
|
// Remove any stale accepts first so a prior unclean exit (e.g. SIGKILL,
|
||||||
|
// where Close never ran) is recovered deterministically rather than
|
||||||
|
// accumulating duplicate rules on the iptables filter table.
|
||||||
|
if err := f.removeAcceptFilterRules(); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("clean stale accept rules: %w", err))
|
||||||
|
}
|
||||||
|
if err := f.openInterface(false); err != nil {
|
||||||
|
merr = multierror.Append(merr, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
a.extMonitor.start()
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// families returns the configured address families (v4, and v6 when present).
|
||||||
|
func (a *InterfaceAllower) families() []*family {
|
||||||
|
families := []*family{a.family4}
|
||||||
|
if a.family6 != nil {
|
||||||
|
families = append(families, a.family6)
|
||||||
|
}
|
||||||
|
return families
|
||||||
|
}
|
||||||
|
|
||||||
|
// reconcileExternalChains re-applies the INPUT accepts to external chains. It
|
||||||
|
// implements externalChainReconciler for the monitor.
|
||||||
|
func (a *InterfaceAllower) reconcileExternalChains() error {
|
||||||
|
var merr *multierror.Error
|
||||||
|
for _, f := range a.families() {
|
||||||
|
if err := f.acceptExternalChainsRules(false); err != nil {
|
||||||
|
merr = multierror.Append(merr, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close stops the monitor and removes the accept rules.
|
||||||
|
func (a *InterfaceAllower) Close() error {
|
||||||
|
a.extMonitor.stop()
|
||||||
|
|
||||||
|
var merr *multierror.Error
|
||||||
|
for _, f := range a.families() {
|
||||||
|
if err := f.removeAcceptFilterRules(); err != nil {
|
||||||
|
merr = multierror.Append(merr, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
@@ -0,0 +1,210 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/google/nftables"
|
||||||
|
"github.com/google/nftables/expr"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/routemanager/refcounter"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (r *family) getIpSet(set firewall.Set, prefixes []netip.Prefix, isSource bool) ([]expr.Any, error) {
|
||||||
|
ref, err := r.ipsetCounter.Increment(set.HashedName(), setInput{
|
||||||
|
set: set,
|
||||||
|
prefixes: prefixes,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("create or get ipset: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return r.getIpSetExprs(ref, isSource)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) createIpSet(setName string, input setInput) (*nftables.Set, error) {
|
||||||
|
// overlapping prefixes will result in an error, so we need to merge them
|
||||||
|
prefixes := firewall.MergeIPRanges(input.prefixes)
|
||||||
|
|
||||||
|
nfset := &nftables.Set{
|
||||||
|
Name: setName,
|
||||||
|
Comment: input.set.Comment(),
|
||||||
|
Table: r.workTable,
|
||||||
|
// required for prefixes
|
||||||
|
Interval: true,
|
||||||
|
KeyType: r.af.setKeyType,
|
||||||
|
}
|
||||||
|
|
||||||
|
elements := r.convertPrefixesToSet(prefixes)
|
||||||
|
nElements := len(elements)
|
||||||
|
|
||||||
|
maxElements := maxPrefixesSet * 2
|
||||||
|
initialElements := elements[:min(maxElements, nElements)]
|
||||||
|
|
||||||
|
if err := r.conn.AddSet(nfset, initialElements); err != nil {
|
||||||
|
return nil, fmt.Errorf("error adding set %s: %w", setName, err)
|
||||||
|
}
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return nil, fmt.Errorf("flush error: %w", err)
|
||||||
|
}
|
||||||
|
log.Debugf("Created new ipset: %s with %d initial prefixes (total prefixes %d)", setName, len(initialElements)/2, len(prefixes))
|
||||||
|
|
||||||
|
// The set is committed now. If a later batch fails, destroy it: the
|
||||||
|
// refcounter records nothing on a create-callback error, so it would
|
||||||
|
// otherwise leak, and a partial source set fails-open for deny rules.
|
||||||
|
if err := r.addRemainingElements(nfset, elements, maxElements); err != nil {
|
||||||
|
if derr := r.deleteIpSet(setName, nfset); derr != nil {
|
||||||
|
log.Warnf("rollback ipset %s after add failure: %v", setName, derr)
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Infof("Created new ipset: %s with %d prefixes", setName, len(prefixes))
|
||||||
|
return nfset, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// addRemainingElements adds element batches beyond the initial one in
|
||||||
|
// maxElements-sized chunks, flushing each. Called after the set has been
|
||||||
|
// created with its first batch.
|
||||||
|
func (r *family) addRemainingElements(nfset *nftables.Set, elements []nftables.SetElement, maxElements int) error {
|
||||||
|
nElements := len(elements)
|
||||||
|
for subStart := maxElements; subStart < nElements; subStart += maxElements {
|
||||||
|
subEnd := min(subStart+maxElements, nElements)
|
||||||
|
subElement := elements[subStart:subEnd]
|
||||||
|
nSubPrefixes := len(subElement) / 2
|
||||||
|
log.Tracef("Adding new prefixes (%d) in ipset: %s", nSubPrefixes, nfset.Name)
|
||||||
|
if err := r.conn.SetAddElements(nfset, subElement); err != nil {
|
||||||
|
return fmt.Errorf("error adding prefixes (%d) to set %s: %w", nSubPrefixes, nfset.Name, err)
|
||||||
|
}
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf("flush error: %w", err)
|
||||||
|
}
|
||||||
|
log.Debugf("Added new prefixes (%d) in ipset: %s", nSubPrefixes, nfset.Name)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) convertPrefixesToSet(prefixes []netip.Prefix) []nftables.SetElement {
|
||||||
|
var elements []nftables.SetElement
|
||||||
|
for _, prefix := range prefixes {
|
||||||
|
// nftables needs half-open intervals [firstIP, lastIP) for prefixes
|
||||||
|
// e.g. 10.0.0.0/24 becomes [10.0.0.0, 10.0.1.0), 10.1.1.1/32 becomes [10.1.1.1, 10.1.1.2) etc
|
||||||
|
firstIP := prefix.Addr()
|
||||||
|
|
||||||
|
// For a /0 the last address is the broadcast and its Next() overflows
|
||||||
|
// to an invalid Addr with an empty key, so wrap to the zero address,
|
||||||
|
// which nftables reads as the open end of a full-range interval.
|
||||||
|
var lastKey []byte
|
||||||
|
if prefix.Bits() == 0 {
|
||||||
|
lastKey = make([]byte, r.af.addrLen)
|
||||||
|
} else {
|
||||||
|
lastKey = calculateLastIP(prefix).Next().AsSlice()
|
||||||
|
}
|
||||||
|
|
||||||
|
// the nft tool also adds a zero-address IntervalEnd element, see https://github.com/google/nftables/issues/247
|
||||||
|
// nftables.SetElement{Key: make([]byte, r.af.addrLen), IntervalEnd: true},
|
||||||
|
elements = append(elements,
|
||||||
|
nftables.SetElement{Key: firstIP.AsSlice()},
|
||||||
|
nftables.SetElement{Key: lastKey, IntervalEnd: true},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return elements
|
||||||
|
}
|
||||||
|
|
||||||
|
// calculateLastIP determines the last IP in a given prefix.
|
||||||
|
func calculateLastIP(prefix netip.Prefix) netip.Addr {
|
||||||
|
masked := prefix.Masked()
|
||||||
|
if masked.Addr().Is4() {
|
||||||
|
hostMask := ^uint32(0) >> masked.Bits()
|
||||||
|
lastIP := uint32FromNetipAddr(masked.Addr()) | hostMask
|
||||||
|
return netip.AddrFrom4(uint32ToBytes(lastIP))
|
||||||
|
}
|
||||||
|
|
||||||
|
// IPv6: set host bits to all 1s
|
||||||
|
b := masked.Addr().As16()
|
||||||
|
bits := masked.Bits()
|
||||||
|
for i := bits; i < 128; i++ {
|
||||||
|
b[i/8] |= 1 << (7 - i%8)
|
||||||
|
}
|
||||||
|
return netip.AddrFrom16(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Utility function to convert netip.Addr to uint32.
|
||||||
|
func uint32FromNetipAddr(addr netip.Addr) uint32 {
|
||||||
|
b := addr.As4()
|
||||||
|
return binary.BigEndian.Uint32(b[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
// Utility function to convert uint32 to a netip-compatible byte slice.
|
||||||
|
func uint32ToBytes(ip uint32) [4]byte {
|
||||||
|
var b [4]byte
|
||||||
|
binary.BigEndian.PutUint32(b[:], ip)
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) deleteIpSet(setName string, nfset *nftables.Set) error {
|
||||||
|
r.conn.DelSet(nfset)
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf(flushError, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugf("Deleted unused ipset %s", setName)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
|
||||||
|
nfset, err := r.conn.GetSetByName(r.workTable, set.HashedName())
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("get set %s: %w", set.HashedName(), err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Overlapping prefixes (e.g. duplicate resolved addresses) make the
|
||||||
|
// interval set reject the batch, so merge them as createIpSet does.
|
||||||
|
prefixes = firewall.MergeIPRanges(prefixes)
|
||||||
|
elements := r.convertPrefixesToSet(prefixes)
|
||||||
|
|
||||||
|
// Add in batches sized like createIpSet so a large update does not
|
||||||
|
// exceed the netlink message size limit.
|
||||||
|
maxElements := maxPrefixesSet * 2
|
||||||
|
for start := 0; start < len(elements); start += maxElements {
|
||||||
|
end := min(start+maxElements, len(elements))
|
||||||
|
if err := r.conn.SetAddElements(nfset, elements[start:end]); err != nil {
|
||||||
|
return fmt.Errorf("add elements to set %s: %w", set.HashedName(), err)
|
||||||
|
}
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
return fmt.Errorf(flushError, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugf("updated set %s with %d prefixes", set.HashedName(), len(prefixes))
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) getIpSetExprs(ref refcounter.Ref[*nftables.Set], isSource bool) ([]expr.Any, error) {
|
||||||
|
// dst offset by default
|
||||||
|
offset := r.af.dstAddrOffset
|
||||||
|
if isSource {
|
||||||
|
// src offset
|
||||||
|
offset = r.af.srcAddrOffset
|
||||||
|
}
|
||||||
|
|
||||||
|
return []expr.Any{
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: offset,
|
||||||
|
Len: r.af.addrLen,
|
||||||
|
},
|
||||||
|
&expr.Lookup{
|
||||||
|
SourceRegister: 1,
|
||||||
|
SetName: ref.Out.Name,
|
||||||
|
SetID: ref.Out.ID,
|
||||||
|
},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,36 @@
|
|||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestConvertPrefixesToSetWildcard verifies that a /0 prefix produces a
|
||||||
|
// usable interval. The last address of a /0 is the broadcast, whose Next()
|
||||||
|
// overflows to an invalid Addr with an empty key; the IntervalEnd must wrap
|
||||||
|
// to the zero address instead so nftables sees a full-range interval.
|
||||||
|
func TestConvertPrefixesToSetWildcard(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
af addrFamily
|
||||||
|
prefix string
|
||||||
|
}{
|
||||||
|
{"IPv4 /0", afIPv4, "0.0.0.0/0"},
|
||||||
|
{"IPv6 /0", afIPv6, "::/0"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
r := &family{af: tt.af}
|
||||||
|
elements := r.convertPrefixesToSet([]netip.Prefix{netip.MustParsePrefix(tt.prefix)})
|
||||||
|
|
||||||
|
require.Len(t, elements, 2, "expected start and interval-end element")
|
||||||
|
assert.False(t, elements[0].IntervalEnd, "first element is the interval start")
|
||||||
|
assert.True(t, elements[1].IntervalEnd, "second element is the interval end")
|
||||||
|
assert.Len(t, elements[1].Key, int(tt.af.addrLen), "interval-end key must be a zero address, not empty")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,85 +0,0 @@
|
|||||||
package nftables
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net"
|
|
||||||
)
|
|
||||||
|
|
||||||
type ipsetStore struct {
|
|
||||||
ipsetReference map[string]int
|
|
||||||
ipsets map[string]map[string]struct{} // ipsetName -> list of ips
|
|
||||||
}
|
|
||||||
|
|
||||||
func newIpsetStore() *ipsetStore {
|
|
||||||
return &ipsetStore{
|
|
||||||
ipsetReference: make(map[string]int),
|
|
||||||
ipsets: make(map[string]map[string]struct{}),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) ips(ipsetName string) (map[string]struct{}, bool) {
|
|
||||||
r, ok := s.ipsets[ipsetName]
|
|
||||||
return r, ok
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) newIpset(ipsetName string) map[string]struct{} {
|
|
||||||
s.ipsetReference[ipsetName] = 0
|
|
||||||
ipList := make(map[string]struct{})
|
|
||||||
s.ipsets[ipsetName] = ipList
|
|
||||||
return ipList
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) deleteIpset(ipsetName string) {
|
|
||||||
delete(s.ipsetReference, ipsetName)
|
|
||||||
delete(s.ipsets, ipsetName)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) DeleteIpFromSet(ipsetName string, ip net.IP) {
|
|
||||||
ipList, ok := s.ipsets[ipsetName]
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
delete(ipList, ip.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) AddIpToSet(ipsetName string, ip net.IP) {
|
|
||||||
ipList, ok := s.ipsets[ipsetName]
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
ipList[ip.String()] = struct{}{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) IsIpInSet(ipsetName string, ip net.IP) bool {
|
|
||||||
ipList, ok := s.ipsets[ipsetName]
|
|
||||||
if !ok {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
_, ok = ipList[ip.String()]
|
|
||||||
return ok
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) AddReferenceToIpset(ipsetName string) {
|
|
||||||
s.ipsetReference[ipsetName]++
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) DeleteReferenceFromIpSet(ipsetName string) {
|
|
||||||
r, ok := s.ipsetReference[ipsetName]
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if r == 0 {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
s.ipsetReference[ipsetName]--
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *ipsetStore) HasReferenceToSet(ipsetName string) bool {
|
|
||||||
if _, ok := s.ipsetReference[ipsetName]; !ok {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if s.ipsetReference[ipsetName] == 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
@@ -3,7 +3,6 @@ package nftables
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -13,10 +12,8 @@ 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"
|
||||||
"github.com/netbirdio/netbird/client/firewall/firewalld"
|
|
||||||
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
||||||
"github.com/netbirdio/netbird/client/internal/statemanager"
|
"github.com/netbirdio/netbird/client/internal/statemanager"
|
||||||
@@ -45,21 +42,17 @@ type iFaceMapper interface {
|
|||||||
Address() wgaddr.Address
|
Address() wgaddr.Address
|
||||||
}
|
}
|
||||||
|
|
||||||
// Manager of iptables firewall
|
// Manager of nftables firewall. Per-family state (peer ACLs, route
|
||||||
|
// ACLs, NAT, DNAT, MSS clamping) lives on family; Manager dispatches
|
||||||
|
// by family and provides the public firewall.Manager surface.
|
||||||
type Manager struct {
|
type Manager struct {
|
||||||
mutex sync.Mutex
|
mutex sync.Mutex
|
||||||
rConn *nftables.Conn
|
rConn *nftables.Conn
|
||||||
wgIface iFaceMapper
|
wgIface iFaceMapper
|
||||||
|
|
||||||
router *router
|
family4 *family
|
||||||
aclManager *AclManager
|
// IPv6 counterpart, nil when no v6 overlay.
|
||||||
|
family6 *family
|
||||||
// IPv6 counterparts, nil when no v6 overlay
|
|
||||||
router6 *router
|
|
||||||
aclManager6 *AclManager
|
|
||||||
|
|
||||||
notrackOutputChain *nftables.Chain
|
|
||||||
notrackPreroutingChain *nftables.Chain
|
|
||||||
|
|
||||||
extMonitor *externalChainMonitor
|
extMonitor *externalChainMonitor
|
||||||
}
|
}
|
||||||
@@ -74,21 +67,10 @@ func Create(wgIface iFaceMapper, mtu uint16) (*Manager, error) {
|
|||||||
tableName := getTableName()
|
tableName := getTableName()
|
||||||
workTable := &nftables.Table{Name: tableName, Family: nftables.TableFamilyIPv4}
|
workTable := &nftables.Table{Name: tableName, Family: nftables.TableFamilyIPv4}
|
||||||
|
|
||||||
var err error
|
m.family4 = newFamily(workTable, wgIface, mtu)
|
||||||
m.router, err = newRouter(workTable, wgIface, mtu)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("create router: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.aclManager, err = newAclManager(workTable, wgIface, chainNameRoutingFw)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("create acl manager: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if wgIface.Address().HasIPv6() {
|
if wgIface.Address().HasIPv6() {
|
||||||
if err := m.createIPv6Components(tableName, wgIface, mtu); err != nil {
|
m.createIPv6Components(tableName, wgIface, mtu)
|
||||||
return nil, fmt.Errorf("create IPv6 firewall: %w", err)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
m.extMonitor = newExternalChainMonitor(m)
|
m.extMonitor = newExternalChainMonitor(m)
|
||||||
@@ -96,30 +78,19 @@ func Create(wgIface iFaceMapper, mtu uint16) (*Manager, error) {
|
|||||||
return m, nil
|
return m, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *Manager) createIPv6Components(tableName string, wgIface iFaceMapper, mtu uint16) error {
|
func (m *Manager) createIPv6Components(tableName string, wgIface iFaceMapper, mtu uint16) {
|
||||||
workTable6 := &nftables.Table{Name: tableName, Family: nftables.TableFamilyIPv6}
|
workTable6 := &nftables.Table{Name: tableName, Family: nftables.TableFamilyIPv6}
|
||||||
|
|
||||||
var err error
|
m.family6 = newFamily(workTable6, wgIface, mtu)
|
||||||
m.router6, err = newRouter(workTable6, wgIface, mtu)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("create v6 router: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Share the per-family forwarding refcounter with the v4 router so a v4
|
// Share the per-family forwarding refcounter with the v4 family so a v4
|
||||||
// rule and a v6 rule against the same state machine cooperate cleanly.
|
// rule and a v6 rule against the same state machine cooperate cleanly.
|
||||||
m.router6.ipFwdState = m.router.ipFwdState
|
m.family6.ipFwdState = m.family4.ipFwdState
|
||||||
|
|
||||||
m.aclManager6, err = newAclManager(workTable6, wgIface, chainNameRoutingFw)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("create v6 acl manager: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// hasIPv6 reports whether the manager has IPv6 components initialized.
|
// hasIPv6 reports whether the manager has IPv6 components initialized.
|
||||||
func (m *Manager) hasIPv6() bool {
|
func (m *Manager) hasIPv6() bool {
|
||||||
return m.router6 != nil
|
return m.family6 != nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *Manager) initIPv6() error {
|
func (m *Manager) initIPv6() error {
|
||||||
@@ -128,12 +99,8 @@ func (m *Manager) initIPv6() error {
|
|||||||
return fmt.Errorf("create v6 work table: %w", err)
|
return fmt.Errorf("create v6 work table: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := m.router6.init(workTable6); err != nil {
|
if err := m.family6.init(workTable6); err != nil {
|
||||||
return fmt.Errorf("v6 router init: %w", err)
|
return fmt.Errorf("v6 family init: %w", err)
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.aclManager6.init(workTable6); err != nil {
|
|
||||||
return fmt.Errorf("v6 acl manager init: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
@@ -156,19 +123,20 @@ func (m *Manager) Init(stateManager *statemanager.Manager) error {
|
|||||||
|
|
||||||
// reconcileExternalChains re-applies passthrough accept rules to external
|
// reconcileExternalChains re-applies passthrough accept rules to external
|
||||||
// filter chains for both IPv4 and IPv6 routers. Called by the monitor when
|
// filter chains for both IPv4 and IPv6 routers. Called by the monitor when
|
||||||
// tables or chains appear (e.g. after firewalld reloads).
|
// tables or chains appear (e.g. after firewalld reloads). Kernel routing opens
|
||||||
|
// both INPUT and FORWARD.
|
||||||
func (m *Manager) reconcileExternalChains() error {
|
func (m *Manager) reconcileExternalChains() error {
|
||||||
m.mutex.Lock()
|
m.mutex.Lock()
|
||||||
defer m.mutex.Unlock()
|
defer m.mutex.Unlock()
|
||||||
|
|
||||||
var merr *multierror.Error
|
var merr *multierror.Error
|
||||||
if m.router != nil {
|
if m.family4 != nil {
|
||||||
if err := m.router.acceptExternalChainsRules(); err != nil {
|
if err := m.family4.acceptExternalChainsRules(true); err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("v4: %w", err))
|
merr = multierror.Append(merr, fmt.Errorf("v4: %w", err))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if m.hasIPv6() {
|
if m.hasIPv6() {
|
||||||
if err := m.router6.acceptExternalChainsRules(); err != nil {
|
if err := m.family6.acceptExternalChainsRules(true); err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("v6: %w", err))
|
merr = multierror.Append(merr, fmt.Errorf("v6: %w", err))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -187,12 +155,8 @@ func (m *Manager) initFirewall() (err error) {
|
|||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
if err := m.router.init(workTable); err != nil {
|
if err := m.family4.init(workTable); err != nil {
|
||||||
return fmt.Errorf("router init: %w", err)
|
return fmt.Errorf("family init: %w", err)
|
||||||
}
|
|
||||||
|
|
||||||
if err := m.aclManager.init(workTable); err != nil {
|
|
||||||
return fmt.Errorf("acl manager init: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if m.hasIPv6() {
|
if m.hasIPv6() {
|
||||||
@@ -202,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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -220,7 +180,7 @@ func (m *Manager) persistState(stateManager *statemanager.Manager) {
|
|||||||
InterfaceState: &InterfaceState{
|
InterfaceState: &InterfaceState{
|
||||||
NameStr: m.wgIface.Name(),
|
NameStr: m.wgIface.Name(),
|
||||||
WGAddress: m.wgIface.Address(),
|
WGAddress: m.wgIface.Address(),
|
||||||
MTU: m.router.mtu,
|
MTU: m.family4.mtu,
|
||||||
},
|
},
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
log.Errorf("failed to update state: %v", err)
|
log.Errorf("failed to update state: %v", err)
|
||||||
@@ -235,12 +195,12 @@ func (m *Manager) persistState(stateManager *statemanager.Manager) {
|
|||||||
|
|
||||||
// rollbackInit performs best-effort cleanup of already-initialized state when Init fails partway through.
|
// rollbackInit performs best-effort cleanup of already-initialized state when Init fails partway through.
|
||||||
func (m *Manager) rollbackInit() {
|
func (m *Manager) rollbackInit() {
|
||||||
if err := m.router.Reset(); err != nil {
|
if err := m.family4.Reset(); err != nil {
|
||||||
log.Warnf("rollback router: %v", err)
|
log.Warnf("rollback family: %v", err)
|
||||||
}
|
}
|
||||||
if m.hasIPv6() {
|
if m.hasIPv6() {
|
||||||
if err := m.router6.Reset(); err != nil {
|
if err := m.family6.Reset(); err != nil {
|
||||||
log.Warnf("rollback v6 router: %v", err)
|
log.Warnf("rollback v6 family: %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := m.cleanupNetbirdTables(); err != nil {
|
if err := m.cleanupNetbirdTables(); err != nil {
|
||||||
@@ -251,118 +211,82 @@ func (m *Manager) rollbackInit() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddPeerFiltering rule to the firewall
|
// AddFilterRule installs a packet-filtering rule.
|
||||||
//
|
//
|
||||||
// If comment argument is empty firewall manager should set
|
// Destination semantics: zero Network → input chain (peer ACL);
|
||||||
// rule ID as comment for the rule
|
// set Network → forward chain (route ACL).
|
||||||
func (m *Manager) AddPeerFiltering(
|
//
|
||||||
id []byte,
|
// Sources are a single address family; the rule is dispatched to the
|
||||||
ip net.IP,
|
// matching per-family backend.
|
||||||
proto firewall.Protocol,
|
func (m *Manager) AddFilterRule(
|
||||||
sPort *firewall.Port,
|
|
||||||
dPort *firewall.Port,
|
|
||||||
action firewall.Action,
|
|
||||||
ipsetName string,
|
|
||||||
) ([]firewall.Rule, error) {
|
|
||||||
m.mutex.Lock()
|
|
||||||
defer m.mutex.Unlock()
|
|
||||||
|
|
||||||
if ip.To4() != nil {
|
|
||||||
return m.aclManager.AddPeerFiltering(id, ip, proto, sPort, dPort, action, ipsetName)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !m.hasIPv6() {
|
|
||||||
return nil, fmt.Errorf("add peer filtering for %s: %w", ip, firewall.ErrIPv6NotInitialized)
|
|
||||||
}
|
|
||||||
return m.aclManager6.AddPeerFiltering(id, ip, proto, sPort, dPort, action, ipsetName)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *Manager) AddRouteFiltering(
|
|
||||||
id []byte,
|
id []byte,
|
||||||
sources []netip.Prefix,
|
sources []netip.Prefix,
|
||||||
destination firewall.Network,
|
destination firewall.Network,
|
||||||
proto firewall.Protocol,
|
proto firewall.Protocol,
|
||||||
sPort, dPort *firewall.Port,
|
sPort *firewall.Port,
|
||||||
|
dPort *firewall.Port,
|
||||||
action firewall.Action,
|
action firewall.Action,
|
||||||
) (firewall.Rule, error) {
|
) (firewall.Rule, error) {
|
||||||
|
if len(sources) == 0 {
|
||||||
|
return nil, firewall.ErrNoSources
|
||||||
|
}
|
||||||
|
|
||||||
m.mutex.Lock()
|
m.mutex.Lock()
|
||||||
defer m.mutex.Unlock()
|
defer m.mutex.Unlock()
|
||||||
|
|
||||||
if isIPv6RouteRule(sources, destination) {
|
fam := m.family4
|
||||||
|
if isIPv6Rule(sources, destination) {
|
||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return nil, fmt.Errorf("add route filtering: %w", firewall.ErrIPv6NotInitialized)
|
return nil, fmt.Errorf("add filtering: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.AddRouteFiltering(id, sources, destination, proto, sPort, dPort, action)
|
fam = m.family6
|
||||||
}
|
}
|
||||||
|
return fam.AddFilterRule(id, sources, destination, proto, sPort, dPort, action)
|
||||||
return m.router.AddRouteFiltering(id, sources, destination, proto, sPort, dPort, action)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// DeletePeerRule from the firewall by rule definition
|
// DeleteFilterRule removes a filtering rule. The owning family is found
|
||||||
func (m *Manager) DeletePeerRule(rule firewall.Rule) error {
|
// by id in the in-memory filter maps, which are the only tracking for
|
||||||
|
// filter rules. family.DeleteFilterRule is idempotent when the id is
|
||||||
|
// absent.
|
||||||
|
func (m *Manager) DeleteFilterRule(rule firewall.Rule) error {
|
||||||
m.mutex.Lock()
|
m.mutex.Lock()
|
||||||
defer m.mutex.Unlock()
|
defer m.mutex.Unlock()
|
||||||
|
|
||||||
if m.hasIPv6() && isIPv6Rule(rule) {
|
fam, err := m.familyForRuleID(rule.ID(), (*family).hasRule, false)
|
||||||
return m.aclManager6.DeletePeerRule(rule)
|
|
||||||
}
|
|
||||||
return m.aclManager.DeletePeerRule(rule)
|
|
||||||
}
|
|
||||||
|
|
||||||
func isIPv6Rule(rule firewall.Rule) bool {
|
|
||||||
r, ok := rule.(*Rule)
|
|
||||||
return ok && r.nftRule != nil && r.nftRule.Table != nil && r.nftRule.Table.Family == nftables.TableFamilyIPv6
|
|
||||||
}
|
|
||||||
|
|
||||||
// isIPv6RouteRule determines whether a route rule belongs to the v6 table.
|
|
||||||
// For static routes, the destination prefix determines the family. For dynamic
|
|
||||||
// routes (DomainSet), the sources determine the family since management
|
|
||||||
// duplicates dynamic rules per family.
|
|
||||||
func isIPv6RouteRule(sources []netip.Prefix, destination firewall.Network) bool {
|
|
||||||
if destination.IsPrefix() {
|
|
||||||
return destination.Prefix.Addr().Is6()
|
|
||||||
}
|
|
||||||
return len(sources) > 0 && sources[0].Addr().Is6()
|
|
||||||
}
|
|
||||||
|
|
||||||
// DeleteRouteRule deletes a routing rule. Route rules live in exactly one
|
|
||||||
// router; the cached maps are normally authoritative, so the kernel is only
|
|
||||||
// consulted when neither map knows about the rule.
|
|
||||||
func (m *Manager) DeleteRouteRule(rule firewall.Rule) error {
|
|
||||||
m.mutex.Lock()
|
|
||||||
defer m.mutex.Unlock()
|
|
||||||
|
|
||||||
id := rule.ID()
|
|
||||||
r, err := m.routerForRuleID(id, (*router).hasRule)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return r.DeleteRouteRule(rule)
|
return fam.DeleteFilterRule(rule)
|
||||||
}
|
}
|
||||||
|
|
||||||
// routerForRuleID picks the router holding the rule with the given id, using
|
// familyForRuleID picks the family holding the rule with the given id, using
|
||||||
// the supplied lookup. If the cached maps disagree (or both miss), it refreshes
|
// the supplied lookup. With refresh set, a miss in both cached maps reloads
|
||||||
// from the kernel once and re-checks before falling back to the v4 router.
|
// the NAT/DNAT rule maps from the kernel once and re-checks before falling
|
||||||
func (m *Manager) routerForRuleID(id string, has func(*router, string) bool) (*router, error) {
|
// back to the v4 family. Filter rules are tracked only in memory and have no
|
||||||
if has(m.router, id) {
|
// kernel-backed reload, so their callers pass refresh as false.
|
||||||
return m.router, nil
|
func (m *Manager) familyForRuleID(id firewall.RuleID, has func(*family, firewall.RuleID) bool, refresh bool) (*family, error) {
|
||||||
}
|
if has(m.family4, id) {
|
||||||
if m.hasIPv6() && has(m.router6, id) {
|
return m.family4, nil
|
||||||
return m.router6, nil
|
|
||||||
}
|
}
|
||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return m.router, nil
|
return m.family4, nil
|
||||||
}
|
}
|
||||||
if err := m.router.refreshRulesMap(); err != nil {
|
if has(m.family6, id) {
|
||||||
|
return m.family6, nil
|
||||||
|
}
|
||||||
|
if !refresh {
|
||||||
|
return m.family4, nil
|
||||||
|
}
|
||||||
|
if err := m.family4.refreshRulesMap(); err != nil {
|
||||||
return nil, fmt.Errorf("refresh v4 rules: %w", err)
|
return nil, fmt.Errorf("refresh v4 rules: %w", err)
|
||||||
}
|
}
|
||||||
if err := m.router6.refreshRulesMap(); err != nil {
|
if err := m.family6.refreshRulesMap(); err != nil {
|
||||||
return nil, fmt.Errorf("refresh v6 rules: %w", err)
|
return nil, fmt.Errorf("refresh v6 rules: %w", err)
|
||||||
}
|
}
|
||||||
if has(m.router6, id) && !has(m.router, id) {
|
if has(m.family6, id) && !has(m.family4, id) {
|
||||||
return m.router6, nil
|
return m.family6, nil
|
||||||
}
|
}
|
||||||
return m.router, nil
|
return m.family4, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *Manager) IsServerRouteSupported() bool {
|
func (m *Manager) IsServerRouteSupported() bool {
|
||||||
@@ -381,10 +305,10 @@ func (m *Manager) AddNatRule(pair firewall.RouterPair) error {
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return fmt.Errorf("add NAT rule: %w", firewall.ErrIPv6NotInitialized)
|
return fmt.Errorf("add NAT rule: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.AddNatRule(pair)
|
return m.family6.AddNatRule(pair)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := m.router.AddNatRule(pair); err != nil {
|
if err := m.family4.AddNatRule(pair); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -396,7 +320,7 @@ func (m *Manager) AddNatRule(pair firewall.RouterPair) error {
|
|||||||
// so the eventual cleanup still works.
|
// so the eventual cleanup still works.
|
||||||
if m.hasIPv6() && pair.Dynamic {
|
if m.hasIPv6() && pair.Dynamic {
|
||||||
v6Pair := firewall.ToV6NatPair(pair)
|
v6Pair := firewall.ToV6NatPair(pair)
|
||||||
if err := m.router6.AddNatRule(v6Pair); err != nil {
|
if err := m.family6.AddNatRule(v6Pair); err != nil {
|
||||||
return fmt.Errorf("add v6 NAT rule: %w", err)
|
return fmt.Errorf("add v6 NAT rule: %w", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -412,18 +336,18 @@ func (m *Manager) RemoveNatRule(pair firewall.RouterPair) error {
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return m.router6.RemoveNatRule(pair)
|
return m.family6.RemoveNatRule(pair)
|
||||||
}
|
}
|
||||||
|
|
||||||
var merr *multierror.Error
|
var merr *multierror.Error
|
||||||
|
|
||||||
if err := m.router.RemoveNatRule(pair); err != nil {
|
if err := m.family4.RemoveNatRule(pair); err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("remove v4 NAT rule: %w", err))
|
merr = multierror.Append(merr, fmt.Errorf("remove v4 NAT rule: %w", err))
|
||||||
}
|
}
|
||||||
|
|
||||||
if m.hasIPv6() && pair.Dynamic {
|
if m.hasIPv6() && pair.Dynamic {
|
||||||
v6Pair := firewall.ToV6NatPair(pair)
|
v6Pair := firewall.ToV6NatPair(pair)
|
||||||
if err := m.router6.RemoveNatRule(v6Pair); err != nil {
|
if err := m.family6.RemoveNatRule(v6Pair); err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("remove v6 NAT rule: %w", err))
|
merr = multierror.Append(merr, fmt.Errorf("remove v6 NAT rule: %w", err))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -431,46 +355,13 @@ func (m *Manager) RemoveNatRule(pair firewall.RouterPair) error {
|
|||||||
return nberrors.FormatErrorOrNil(merr)
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
}
|
}
|
||||||
|
|
||||||
// AllowNetbird allows netbird interface traffic.
|
|
||||||
// This is called when USPFilter wraps the native firewall, adding blanket accept
|
|
||||||
// rules so that packet filtering is handled in userspace instead of by netfilter.
|
|
||||||
//
|
|
||||||
// TODO: In USP mode this only adds ACCEPT to the netbird table's own chains,
|
|
||||||
// which doesn't override DROP rules in external tables (e.g. firewalld).
|
|
||||||
// Should add passthrough rules to external chains (like the native mode router's
|
|
||||||
// addExternalChainsRules does) for both the netbird table family and inet tables.
|
|
||||||
// The netbird table itself is fine (routing chains already exist there), but
|
|
||||||
// non-netbird tables with INPUT/FORWARD hooks can still DROP our WG traffic.
|
|
||||||
func (m *Manager) AllowNetbird() error {
|
|
||||||
m.mutex.Lock()
|
|
||||||
defer m.mutex.Unlock()
|
|
||||||
|
|
||||||
if err := m.aclManager.createDefaultAllowRules(); err != nil {
|
|
||||||
return fmt.Errorf("create default allow rules: %w", err)
|
|
||||||
}
|
|
||||||
if m.hasIPv6() {
|
|
||||||
if err := m.aclManager6.createDefaultAllowRules(); err != nil {
|
|
||||||
return fmt.Errorf("create v6 default allow rules: %w", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := m.rConn.Flush(); err != nil {
|
|
||||||
return fmt.Errorf("flush allow input netbird rules: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := firewalld.TrustInterface(m.wgIface.Name()); err != nil {
|
|
||||||
log.Warnf("failed to trust interface in firewalld: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetLegacyManagement sets the route manager to use legacy management
|
// SetLegacyManagement sets the route manager to use legacy management
|
||||||
func (m *Manager) SetLegacyManagement(isLegacy bool) error {
|
func (m *Manager) SetLegacyManagement(isLegacy bool) error {
|
||||||
if err := firewall.SetLegacyManagement(m.router, isLegacy); err != nil {
|
if err := firewall.SetLegacyManagement(m.family4, isLegacy); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if m.hasIPv6() {
|
if m.hasIPv6() {
|
||||||
return firewall.SetLegacyManagement(m.router6, isLegacy)
|
return firewall.SetLegacyManagement(m.family6, isLegacy)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -484,13 +375,13 @@ func (m *Manager) Close(stateManager *statemanager.Manager) error {
|
|||||||
|
|
||||||
var merr *multierror.Error
|
var merr *multierror.Error
|
||||||
|
|
||||||
if err := m.router.Reset(); err != nil {
|
if err := m.family4.Reset(); err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("reset router: %v", err))
|
merr = multierror.Append(merr, fmt.Errorf("reset family: %w", err))
|
||||||
}
|
}
|
||||||
|
|
||||||
if m.hasIPv6() {
|
if m.hasIPv6() {
|
||||||
if err := m.router6.Reset(); err != nil {
|
if err := m.family6.Reset(); err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("reset v6 router: %v", err))
|
merr = multierror.Append(merr, fmt.Errorf("reset v6 family: %w", err))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -531,11 +422,11 @@ func (m *Manager) SetLogLevel(log.Level) {
|
|||||||
|
|
||||||
func (m *Manager) EnableRouting() error {
|
func (m *Manager) EnableRouting() error {
|
||||||
// v6 only when the overlay actually has v6.
|
// v6 only when the overlay actually has v6.
|
||||||
return m.router.ipFwdState.RequestRouting(m.router6 != nil)
|
return m.family4.ipFwdState.RequestRouting(m.hasIPv6())
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *Manager) DisableRouting() error {
|
func (m *Manager) DisableRouting() error {
|
||||||
return m.router.ipFwdState.ReleaseRouting()
|
return m.family4.ipFwdState.ReleaseRouting()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Flush rule/chain/set operations from the buffer
|
// Flush rule/chain/set operations from the buffer
|
||||||
@@ -546,20 +437,16 @@ func (m *Manager) Flush() error {
|
|||||||
m.mutex.Lock()
|
m.mutex.Lock()
|
||||||
defer m.mutex.Unlock()
|
defer m.mutex.Unlock()
|
||||||
|
|
||||||
if err := m.aclManager.Flush(); err != nil {
|
if err := m.family4.Flush(); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
if m.hasIPv6() {
|
if m.hasIPv6() {
|
||||||
if err := m.aclManager6.Flush(); err != nil {
|
if err := m.family6.Flush(); err != nil {
|
||||||
return fmt.Errorf("flush v6 acl: %w", err)
|
return fmt.Errorf("flush v6 family: %w", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := m.refreshNoTrackChains(); err != nil {
|
|
||||||
log.Errorf("failed to refresh notrack chains: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -572,9 +459,9 @@ func (m *Manager) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error)
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return nil, fmt.Errorf("add DNAT rule: %w", firewall.ErrIPv6NotInitialized)
|
return nil, fmt.Errorf("add DNAT rule: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.AddDNATRule(rule)
|
return m.family6.AddDNATRule(rule)
|
||||||
}
|
}
|
||||||
return m.router.AddDNATRule(rule)
|
return m.family4.AddDNATRule(rule)
|
||||||
}
|
}
|
||||||
|
|
||||||
// DeleteDNATRule deletes a DNAT rule
|
// DeleteDNATRule deletes a DNAT rule
|
||||||
@@ -582,7 +469,7 @@ func (m *Manager) DeleteDNATRule(rule firewall.Rule) error {
|
|||||||
m.mutex.Lock()
|
m.mutex.Lock()
|
||||||
defer m.mutex.Unlock()
|
defer m.mutex.Unlock()
|
||||||
|
|
||||||
r, err := m.routerForRuleID(rule.ID(), (*router).hasDNATRule)
|
r, err := m.familyForRuleID(rule.ID(), (*family).hasDNATRule, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -603,12 +490,12 @@ func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := m.router.UpdateSet(set, v4Prefixes); err != nil {
|
if err := m.family4.UpdateSet(set, v4Prefixes); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
if m.hasIPv6() && len(v6Prefixes) > 0 {
|
if m.hasIPv6() && len(v6Prefixes) > 0 {
|
||||||
if err := m.router6.UpdateSet(set, v6Prefixes); err != nil {
|
if err := m.family6.UpdateSet(set, v6Prefixes); err != nil {
|
||||||
return fmt.Errorf("update v6 set: %w", err)
|
return fmt.Errorf("update v6 set: %w", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -625,9 +512,9 @@ func (m *Manager) AddInboundDNAT(localAddr netip.Addr, protocol firewall.Protoco
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return fmt.Errorf("add inbound DNAT: %w", firewall.ErrIPv6NotInitialized)
|
return fmt.Errorf("add inbound DNAT: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.AddInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family6.AddInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
return m.router.AddInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family4.AddInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
|
|
||||||
// RemoveInboundDNAT removes an inbound DNAT rule.
|
// RemoveInboundDNAT removes an inbound DNAT rule.
|
||||||
@@ -639,9 +526,9 @@ func (m *Manager) RemoveInboundDNAT(localAddr netip.Addr, protocol firewall.Prot
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return fmt.Errorf("remove inbound DNAT: %w", firewall.ErrIPv6NotInitialized)
|
return fmt.Errorf("remove inbound DNAT: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.RemoveInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family6.RemoveInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
return m.router.RemoveInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family4.RemoveInboundDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddOutputDNAT adds an OUTPUT chain DNAT rule for locally-generated traffic.
|
// AddOutputDNAT adds an OUTPUT chain DNAT rule for locally-generated traffic.
|
||||||
@@ -653,9 +540,9 @@ func (m *Manager) AddOutputDNAT(localAddr netip.Addr, protocol firewall.Protocol
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return fmt.Errorf("add output DNAT: %w", firewall.ErrIPv6NotInitialized)
|
return fmt.Errorf("add output DNAT: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.AddOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family6.AddOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
return m.router.AddOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family4.AddOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
|
|
||||||
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
|
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
|
||||||
@@ -667,179 +554,9 @@ func (m *Manager) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Proto
|
|||||||
if !m.hasIPv6() {
|
if !m.hasIPv6() {
|
||||||
return fmt.Errorf("remove output DNAT: %w", firewall.ErrIPv6NotInitialized)
|
return fmt.Errorf("remove output DNAT: %w", firewall.ErrIPv6NotInitialized)
|
||||||
}
|
}
|
||||||
return m.router6.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
return m.family6.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||||
}
|
}
|
||||||
return m.router.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) {
|
||||||
@@ -898,3 +615,14 @@ func getEstablishedExprs(register uint32) []expr.Any {
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// isIPv6Rule reports whether the rule belongs to the v6 table. For a
|
||||||
|
// prefix destination the destination family decides; otherwise the
|
||||||
|
// (single-family) sources do, since management duplicates rules per
|
||||||
|
// family.
|
||||||
|
func isIPv6Rule(sources []netip.Prefix, destination firewall.Network) bool {
|
||||||
|
if destination.IsPrefix() {
|
||||||
|
return destination.Prefix.Addr().Is6()
|
||||||
|
}
|
||||||
|
return len(sources) > 0 && sources[0].Addr().Is6()
|
||||||
|
}
|
||||||
|
|||||||
@@ -72,13 +72,13 @@ func TestNftablesManager(t *testing.T) {
|
|||||||
|
|
||||||
testClient := &nftables.Conn{}
|
testClient := &nftables.Conn{}
|
||||||
|
|
||||||
rule, err := manager.AddPeerFiltering(nil, ip.AsSlice(), fw.ProtocolTCP, nil, &fw.Port{Values: []uint16{53}}, fw.ActionDrop, "")
|
rule, err := manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, fw.ProtocolTCP, nil, &fw.Port{Values: []uint16{53}}, fw.ActionDrop)
|
||||||
require.NoError(t, err, "failed to add rule")
|
require.NoError(t, err, "failed to add rule")
|
||||||
|
|
||||||
err = manager.Flush()
|
err = manager.Flush()
|
||||||
require.NoError(t, err, "failed to flush")
|
require.NoError(t, err, "failed to flush")
|
||||||
|
|
||||||
rules, err := testClient.GetRules(manager.aclManager.workTable, manager.aclManager.chainInputRules)
|
rules, err := testClient.GetRules(manager.family4.workTable, manager.family4.chainInputRules)
|
||||||
require.NoError(t, err, "failed to get rules")
|
require.NoError(t, err, "failed to get rules")
|
||||||
|
|
||||||
require.Len(t, rules, 2, "expected 2 rules")
|
require.Len(t, rules, 2, "expected 2 rules")
|
||||||
@@ -149,15 +149,12 @@ func TestNftablesManager(t *testing.T) {
|
|||||||
// Compare connection tracking rule at position 1 (pushed down by DROP rule insertion)
|
// Compare connection tracking rule at position 1 (pushed down by DROP rule insertion)
|
||||||
compareExprsIgnoringCounters(t, rules[1].Exprs, expectedExprs1)
|
compareExprsIgnoringCounters(t, rules[1].Exprs, expectedExprs1)
|
||||||
|
|
||||||
for _, r := range rule {
|
require.NoError(t, manager.DeleteFilterRule(rule), "failed to delete rule")
|
||||||
err = manager.DeletePeerRule(r)
|
|
||||||
require.NoError(t, err, "failed to delete rule")
|
|
||||||
}
|
|
||||||
|
|
||||||
err = manager.Flush()
|
err = manager.Flush()
|
||||||
require.NoError(t, err, "failed to flush")
|
require.NoError(t, err, "failed to flush")
|
||||||
|
|
||||||
rules, err = testClient.GetRules(manager.aclManager.workTable, manager.aclManager.chainInputRules)
|
rules, err = testClient.GetRules(manager.family4.workTable, manager.family4.chainInputRules)
|
||||||
require.NoError(t, err, "failed to get rules")
|
require.NoError(t, err, "failed to get rules")
|
||||||
// established rule remains
|
// established rule remains
|
||||||
require.Len(t, rules, 1, "expected 1 rules after deletion")
|
require.Len(t, rules, 1, "expected 1 rules after deletion")
|
||||||
@@ -182,47 +179,39 @@ func TestNftablesManagerRuleOrder(t *testing.T) {
|
|||||||
testClient := &nftables.Conn{}
|
testClient := &nftables.Conn{}
|
||||||
|
|
||||||
// Add accept rule first
|
// Add accept rule first
|
||||||
_, err = manager.AddPeerFiltering(nil, ip.AsSlice(), fw.ProtocolTCP, nil, &fw.Port{Values: []uint16{80}}, fw.ActionAccept, "accept-http")
|
_, err = manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, fw.ProtocolTCP, nil, &fw.Port{Values: []uint16{80}}, fw.ActionAccept)
|
||||||
require.NoError(t, err, "failed to add accept rule")
|
require.NoError(t, err, "failed to add accept rule")
|
||||||
|
|
||||||
// Add deny rule second for the same traffic
|
// Add deny rule second for the same traffic
|
||||||
_, err = manager.AddPeerFiltering(nil, ip.AsSlice(), fw.ProtocolTCP, nil, &fw.Port{Values: []uint16{80}}, fw.ActionDrop, "deny-http")
|
_, err = manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, fw.ProtocolTCP, nil, &fw.Port{Values: []uint16{80}}, fw.ActionDrop)
|
||||||
require.NoError(t, err, "failed to add deny rule")
|
require.NoError(t, err, "failed to add deny rule")
|
||||||
|
|
||||||
err = manager.Flush()
|
err = manager.Flush()
|
||||||
require.NoError(t, err, "failed to flush")
|
require.NoError(t, err, "failed to flush")
|
||||||
|
|
||||||
rules, err := testClient.GetRules(manager.aclManager.workTable, manager.aclManager.chainInputRules)
|
rules, err := testClient.GetRules(manager.family4.workTable, manager.family4.chainInputRules)
|
||||||
require.NoError(t, err, "failed to get rules")
|
require.NoError(t, err, "failed to get rules")
|
||||||
|
|
||||||
t.Logf("Found %d rules in nftables chain", len(rules))
|
t.Logf("Found %d rules in nftables chain", len(rules))
|
||||||
|
|
||||||
// Find the accept and deny rules and verify deny comes before accept
|
// Single-source rules emit a direct payload+cmp on the source IP
|
||||||
|
// (no set lookup). Match by source-IP + port + verdict instead of
|
||||||
|
// the legacy per-(action,port) set names ("deny-http"/"accept-http")
|
||||||
|
// that this test predates.
|
||||||
|
wantSrc := ip.AsSlice()
|
||||||
var acceptRuleIndex, denyRuleIndex = -1, -1
|
var acceptRuleIndex, denyRuleIndex = -1, -1
|
||||||
for i, rule := range rules {
|
for i, rule := range rules {
|
||||||
hasAcceptHTTPSet := false
|
var hasSrc, hasPort80 bool
|
||||||
hasDenyHTTPSet := false
|
|
||||||
hasPort80 := false
|
|
||||||
var action string
|
var action string
|
||||||
|
|
||||||
for _, e := range rule.Exprs {
|
for _, e := range rule.Exprs {
|
||||||
// Check for set lookup
|
if cmp, ok := e.(*expr.Cmp); ok && cmp.Op == expr.CmpOpEq {
|
||||||
if lookup, ok := e.(*expr.Lookup); ok {
|
if bytes.Equal(cmp.Data, wantSrc) {
|
||||||
switch lookup.SetName {
|
hasSrc = true
|
||||||
case "accept-http":
|
|
||||||
hasAcceptHTTPSet = true
|
|
||||||
case "deny-http":
|
|
||||||
hasDenyHTTPSet = true
|
|
||||||
}
|
}
|
||||||
|
if len(cmp.Data) == 2 && binary.BigEndian.Uint16(cmp.Data) == 80 {
|
||||||
}
|
|
||||||
// Check for port 80
|
|
||||||
if cmp, ok := e.(*expr.Cmp); ok {
|
|
||||||
if cmp.Op == expr.CmpOpEq && len(cmp.Data) == 2 && binary.BigEndian.Uint16(cmp.Data) == 80 {
|
|
||||||
hasPort80 = true
|
hasPort80 = true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// Check for verdict
|
|
||||||
if verdict, ok := e.(*expr.Verdict); ok {
|
if verdict, ok := e.(*expr.Verdict); ok {
|
||||||
switch verdict.Kind {
|
switch verdict.Kind {
|
||||||
case expr.VerdictAccept:
|
case expr.VerdictAccept:
|
||||||
@@ -233,11 +222,15 @@ func TestNftablesManagerRuleOrder(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if hasAcceptHTTPSet && hasPort80 && action == "ACCEPT" {
|
if !hasSrc || !hasPort80 {
|
||||||
t.Logf("Rule [%d]: accept-http set + Port 80 + ACCEPT", i)
|
continue
|
||||||
|
}
|
||||||
|
switch action {
|
||||||
|
case "ACCEPT":
|
||||||
|
t.Logf("Rule [%d]: src=%s port=80 ACCEPT", i, ip)
|
||||||
acceptRuleIndex = i
|
acceptRuleIndex = i
|
||||||
} else if hasDenyHTTPSet && hasPort80 && action == "DROP" {
|
case "DROP":
|
||||||
t.Logf("Rule [%d]: deny-http set + Port 80 + DROP", i)
|
t.Logf("Rule [%d]: src=%s port=80 DROP", i, ip)
|
||||||
denyRuleIndex = i
|
denyRuleIndex = i
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -281,7 +274,7 @@ func TestNFtablesCreatePerformance(t *testing.T) {
|
|||||||
start := time.Now()
|
start := time.Now()
|
||||||
for i := 0; i < testMax; i++ {
|
for i := 0; i < testMax; i++ {
|
||||||
port := &fw.Port{Values: []uint16{uint16(1000 + i)}}
|
port := &fw.Port{Values: []uint16{uint16(1000 + i)}}
|
||||||
_, err = manager.AddPeerFiltering(nil, ip.AsSlice(), "tcp", nil, port, fw.ActionAccept, "")
|
_, err = manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, "tcp", nil, port, fw.ActionAccept)
|
||||||
require.NoError(t, err, "failed to add rule")
|
require.NoError(t, err, "failed to add rule")
|
||||||
|
|
||||||
if i%100 == 0 {
|
if i%100 == 0 {
|
||||||
@@ -363,10 +356,10 @@ func TestNftablesManagerCompatibilityWithIptables(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
ip := netip.MustParseAddr("100.96.0.1")
|
ip := netip.MustParseAddr("100.96.0.1")
|
||||||
_, err = manager.AddPeerFiltering(nil, ip.AsSlice(), fw.ProtocolTCP, nil, &fw.Port{Values: []uint16{80}}, fw.ActionAccept, "")
|
_, err = manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, fw.ProtocolTCP, nil, &fw.Port{Values: []uint16{80}}, fw.ActionAccept)
|
||||||
require.NoError(t, err, "failed to add peer filtering rule")
|
require.NoError(t, err, "failed to add peer filtering rule")
|
||||||
|
|
||||||
_, err = manager.AddRouteFiltering(
|
_, err = manager.AddFilterRule(
|
||||||
nil,
|
nil,
|
||||||
[]netip.Prefix{netip.MustParsePrefix("192.168.2.0/24")},
|
[]netip.Prefix{netip.MustParsePrefix("192.168.2.0/24")},
|
||||||
fw.Network{Prefix: netip.MustParsePrefix("10.1.0.0/24")},
|
fw.Network{Prefix: netip.MustParsePrefix("10.1.0.0/24")},
|
||||||
@@ -439,10 +432,10 @@ func TestNftablesManagerIPv6CompatibilityWithIp6tables(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
ip := netip.MustParseAddr("fd00::2")
|
ip := netip.MustParseAddr("fd00::2")
|
||||||
_, err = manager.AddPeerFiltering(nil, ip.AsSlice(), fw.ProtocolTCP, nil, &fw.Port{Values: []uint16{80}}, fw.ActionAccept, "")
|
_, err = manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, fw.ProtocolTCP, nil, &fw.Port{Values: []uint16{80}}, fw.ActionAccept)
|
||||||
require.NoError(t, err, "add v6 peer filtering rule")
|
require.NoError(t, err, "add v6 peer filtering rule")
|
||||||
|
|
||||||
_, err = manager.AddRouteFiltering(
|
_, err = manager.AddFilterRule(
|
||||||
nil,
|
nil,
|
||||||
[]netip.Prefix{netip.MustParsePrefix("fd00:1::/64")},
|
[]netip.Prefix{netip.MustParsePrefix("fd00:1::/64")},
|
||||||
fw.Network{Prefix: netip.MustParsePrefix("2001:db8::/48")},
|
fw.Network{Prefix: netip.MustParsePrefix("2001:db8::/48")},
|
||||||
@@ -552,7 +545,7 @@ func TestNftablesManagerCompatibilityWithIptablesFor6kPrefixes(t *testing.T) {
|
|||||||
prefixes = append(prefixes, netip.PrefixFrom(addr, 24))
|
prefixes = append(prefixes, netip.PrefixFrom(addr, 24))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
_, err = manager.AddRouteFiltering(
|
_, err = manager.AddFilterRule(
|
||||||
nil,
|
nil,
|
||||||
prefixes,
|
prefixes,
|
||||||
fw.Network{Prefix: netip.MustParsePrefix("10.2.0.0/24")},
|
fw.Network{Prefix: netip.MustParsePrefix("10.2.0.0/24")},
|
||||||
@@ -567,7 +560,7 @@ func TestNftablesManagerCompatibilityWithIptablesFor6kPrefixes(t *testing.T) {
|
|||||||
verifyIptablesOutput(t, stdout, stderr)
|
verifyIptablesOutput(t, stdout, stderr)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestNftablesManagerCompatibilityWithIptablesForEmptyPrefixes(t *testing.T) {
|
func TestNftablesManagerCompatibilityWithIptablesForWildcardSource(t *testing.T) {
|
||||||
if check() != NFTABLES {
|
if check() != NFTABLES {
|
||||||
t.Skip("nftables not supported on this system")
|
t.Skip("nftables not supported on this system")
|
||||||
}
|
}
|
||||||
@@ -593,9 +586,9 @@ func TestNftablesManagerCompatibilityWithIptablesForEmptyPrefixes(t *testing.T)
|
|||||||
verifyIptablesOutput(t, stdout, stderr)
|
verifyIptablesOutput(t, stdout, stderr)
|
||||||
})
|
})
|
||||||
|
|
||||||
_, err = manager.AddRouteFiltering(
|
_, err = manager.AddFilterRule(
|
||||||
nil,
|
nil,
|
||||||
[]netip.Prefix{},
|
[]netip.Prefix{netip.MustParsePrefix("0.0.0.0/0")},
|
||||||
fw.Network{Prefix: netip.MustParsePrefix("10.2.0.0/24")},
|
fw.Network{Prefix: netip.MustParsePrefix("10.2.0.0/24")},
|
||||||
fw.ProtocolTCP,
|
fw.ProtocolTCP,
|
||||||
nil,
|
nil,
|
||||||
@@ -608,6 +601,73 @@ func TestNftablesManagerCompatibilityWithIptablesForEmptyPrefixes(t *testing.T)
|
|||||||
verifyIptablesOutput(t, stdout, stderr)
|
verifyIptablesOutput(t, stdout, stderr)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestNftablesManagerMultiPortFilter(t *testing.T) {
|
||||||
|
if check() != NFTABLES {
|
||||||
|
t.Skip("nftables not supported on this system")
|
||||||
|
}
|
||||||
|
|
||||||
|
manager, err := Create(ifaceMock, iface.DefaultMTU)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, manager.Init(nil))
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
require.NoError(t, manager.Close(nil), "failed to reset manager state")
|
||||||
|
})
|
||||||
|
|
||||||
|
ip := netip.MustParseAddr("100.96.0.1")
|
||||||
|
|
||||||
|
rule, err := manager.AddFilterRule(nil, pfx(ip.AsSlice()), fw.Network{}, fw.ProtocolTCP, nil, &fw.Port{Values: []uint16{80, 443}}, fw.ActionAccept)
|
||||||
|
require.NoError(t, err, "failed to add multi-port rule")
|
||||||
|
|
||||||
|
testClient := &nftables.Conn{}
|
||||||
|
rules, err := testClient.GetRules(manager.family4.workTable, manager.family4.chainInputRules)
|
||||||
|
require.NoError(t, err, "failed to get rules")
|
||||||
|
|
||||||
|
var lookup *expr.Lookup
|
||||||
|
for _, kernelRule := range rules {
|
||||||
|
if string(kernelRule.UserData) != string(rule.ID()) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, e := range kernelRule.Exprs {
|
||||||
|
if l, ok := e.(*expr.Lookup); ok {
|
||||||
|
lookup = l
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
require.NotNil(t, lookup, "multi-port rule must match ports via a set lookup")
|
||||||
|
|
||||||
|
sets, err := testClient.GetSets(manager.family4.workTable)
|
||||||
|
require.NoError(t, err, "failed to get sets")
|
||||||
|
|
||||||
|
var portSet *nftables.Set
|
||||||
|
for _, s := range sets {
|
||||||
|
if s.Name == lookup.SetName {
|
||||||
|
portSet = s
|
||||||
|
}
|
||||||
|
}
|
||||||
|
require.NotNil(t, portSet, "anonymous port set not found in kernel")
|
||||||
|
|
||||||
|
portSet.Table = manager.family4.workTable
|
||||||
|
elements, err := testClient.GetSetElements(portSet)
|
||||||
|
require.NoError(t, err, "failed to get set elements")
|
||||||
|
|
||||||
|
ports := make(map[uint16]bool)
|
||||||
|
for _, e := range elements {
|
||||||
|
require.Len(t, e.Key, 2, "port set element key should be 2 bytes")
|
||||||
|
ports[binary.BigEndian.Uint16(e.Key)] = true
|
||||||
|
}
|
||||||
|
require.True(t, ports[80], "port set should contain port 80")
|
||||||
|
require.True(t, ports[443], "port set should contain port 443")
|
||||||
|
|
||||||
|
require.NoError(t, manager.DeleteFilterRule(rule), "failed to delete rule")
|
||||||
|
|
||||||
|
rules, err = testClient.GetRules(manager.family4.workTable, manager.family4.chainInputRules)
|
||||||
|
require.NoError(t, err, "failed to get rules after delete")
|
||||||
|
for _, kernelRule := range rules {
|
||||||
|
require.NotEqual(t, string(rule.ID()), string(kernelRule.UserData), "rule should be removed from kernel")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func compareExprsIgnoringCounters(t *testing.T, got, want []expr.Any) {
|
func compareExprsIgnoringCounters(t *testing.T, got, want []expr.Any) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
require.Equal(t, len(got), len(want), "expression count mismatch")
|
require.Equal(t, len(got), len(want), "expression count mismatch")
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -37,7 +37,7 @@ func TestNftablesManager_AddNatRule(t *testing.T) {
|
|||||||
|
|
||||||
for _, testCase := range test.InsertRuleTestCases {
|
for _, testCase := range test.InsertRuleTestCases {
|
||||||
t.Run(testCase.Name, func(t *testing.T) {
|
t.Run(testCase.Name, func(t *testing.T) {
|
||||||
// need fw manager to init both acl mgr and router for all chains to be present
|
// need fw manager to init both acl mgr and family for all chains to be present
|
||||||
manager, err := Create(ifaceMock, iface.DefaultMTU)
|
manager, err := Create(ifaceMock, iface.DefaultMTU)
|
||||||
t.Cleanup(func() {
|
t.Cleanup(func() {
|
||||||
require.NoError(t, manager.Close(nil))
|
require.NoError(t, manager.Close(nil))
|
||||||
@@ -47,7 +47,7 @@ func TestNftablesManager_AddNatRule(t *testing.T) {
|
|||||||
|
|
||||||
nftablesTestingClient := &nftables.Conn{}
|
nftablesTestingClient := &nftables.Conn{}
|
||||||
|
|
||||||
rtr := manager.router
|
rtr := manager.family4
|
||||||
err = rtr.AddNatRule(testCase.InputPair)
|
err = rtr.AddNatRule(testCase.InputPair)
|
||||||
require.NoError(t, err, "pair should be inserted")
|
require.NoError(t, err, "pair should be inserted")
|
||||||
|
|
||||||
@@ -90,9 +90,9 @@ func TestNftablesManager_AddNatRule(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Build CIDR matching expressions
|
// Build CIDR matching expressions
|
||||||
testRouter := &router{af: afIPv4}
|
testRouter := &family{af: afIPv4}
|
||||||
sourceExp := testRouter.applyPrefix(testCase.InputPair.Source.Prefix, true)
|
sourceExp := prefixMatchExprs(testRouter.af, testCase.InputPair.Source.Prefix, true)
|
||||||
destExp := testRouter.applyPrefix(testCase.InputPair.Destination.Prefix, false)
|
destExp := prefixMatchExprs(testRouter.af, testCase.InputPair.Destination.Prefix, false)
|
||||||
|
|
||||||
// Combine all expressions in the correct order
|
// Combine all expressions in the correct order
|
||||||
// nolint:gocritic
|
// nolint:gocritic
|
||||||
@@ -100,14 +100,14 @@ func TestNftablesManager_AddNatRule(t *testing.T) {
|
|||||||
testingExpression = append(testingExpression, sourceExp...)
|
testingExpression = append(testingExpression, sourceExp...)
|
||||||
testingExpression = append(testingExpression, destExp...)
|
testingExpression = append(testingExpression, destExp...)
|
||||||
|
|
||||||
natRuleKey := firewall.GenKey(firewall.PreroutingFormat, testCase.InputPair)
|
natRuleKey := testCase.InputPair.GenKey(firewall.PreroutingFormat)
|
||||||
found := 0
|
found := 0
|
||||||
for _, chain := range rtr.chains {
|
for _, chain := range rtr.chains {
|
||||||
if chain.Name == chainNameManglePrerouting {
|
if chain.Name == chainNameManglePrerouting {
|
||||||
rules, err := nftablesTestingClient.GetRules(chain.Table, chain)
|
rules, err := nftablesTestingClient.GetRules(chain.Table, chain)
|
||||||
require.NoError(t, err, "should list rules for %s table and %s chain", chain.Table.Name, chain.Name)
|
require.NoError(t, err, "should list rules for %s table and %s chain", chain.Table.Name, chain.Name)
|
||||||
for _, rule := range rules {
|
for _, rule := range rules {
|
||||||
if len(rule.UserData) > 0 && string(rule.UserData) == natRuleKey {
|
if len(rule.UserData) > 0 && firewall.RuleID(rule.UserData) == natRuleKey {
|
||||||
// Compare expressions up to the mark setting expressions
|
// Compare expressions up to the mark setting expressions
|
||||||
require.ElementsMatchf(t, rule.Exprs[:len(testingExpression)], testingExpression, "prerouting nat rule elements should match")
|
require.ElementsMatchf(t, rule.Exprs[:len(testingExpression)], testingExpression, "prerouting nat rule elements should match")
|
||||||
found = 1
|
found = 1
|
||||||
@@ -135,19 +135,19 @@ func TestNftablesManager_RemoveNatRule(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.NoError(t, manager.Init(nil))
|
require.NoError(t, manager.Init(nil))
|
||||||
|
|
||||||
rtr := manager.router
|
rtr := manager.family4
|
||||||
|
|
||||||
// First add the NAT rule using the router's method
|
// First add the NAT rule using the family's method
|
||||||
err = rtr.AddNatRule(testCase.InputPair)
|
err = rtr.AddNatRule(testCase.InputPair)
|
||||||
require.NoError(t, err, "should add NAT rule")
|
require.NoError(t, err, "should add NAT rule")
|
||||||
|
|
||||||
// Verify the rule was added
|
// Verify the rule was added
|
||||||
natRuleKey := firewall.GenKey(firewall.PreroutingFormat, testCase.InputPair)
|
natRuleKey := testCase.InputPair.GenKey(firewall.PreroutingFormat)
|
||||||
found := false
|
found := false
|
||||||
rules, err := rtr.conn.GetRules(rtr.workTable, rtr.chains[chainNameManglePrerouting])
|
rules, err := rtr.conn.GetRules(rtr.workTable, rtr.chains[chainNameManglePrerouting])
|
||||||
require.NoError(t, err, "should list rules")
|
require.NoError(t, err, "should list rules")
|
||||||
for _, rule := range rules {
|
for _, rule := range rules {
|
||||||
if len(rule.UserData) > 0 && string(rule.UserData) == natRuleKey {
|
if len(rule.UserData) > 0 && firewall.RuleID(rule.UserData) == natRuleKey {
|
||||||
found = true
|
found = true
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
@@ -163,7 +163,7 @@ func TestNftablesManager_RemoveNatRule(t *testing.T) {
|
|||||||
rules, err = rtr.conn.GetRules(rtr.workTable, rtr.chains[chainNameManglePrerouting])
|
rules, err = rtr.conn.GetRules(rtr.workTable, rtr.chains[chainNameManglePrerouting])
|
||||||
require.NoError(t, err, "should list rules after removal")
|
require.NoError(t, err, "should list rules after removal")
|
||||||
for _, rule := range rules {
|
for _, rule := range rules {
|
||||||
if len(rule.UserData) > 0 && string(rule.UserData) == natRuleKey {
|
if len(rule.UserData) > 0 && firewall.RuleID(rule.UserData) == natRuleKey {
|
||||||
found = true
|
found = true
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
@@ -200,11 +200,10 @@ func TestRouter_AddRouteFiltering(t *testing.T) {
|
|||||||
|
|
||||||
defer deleteWorkTable()
|
defer deleteWorkTable()
|
||||||
|
|
||||||
r, err := newRouter(workTable, ifaceMock, iface.DefaultMTU)
|
r := newFamily(workTable, ifaceMock, iface.DefaultMTU)
|
||||||
require.NoError(t, err, "Failed to create router")
|
|
||||||
require.NoError(t, r.init(workTable))
|
require.NoError(t, r.init(workTable))
|
||||||
|
|
||||||
defer func(r *router) {
|
defer func(r *family) {
|
||||||
require.NoError(t, r.Reset(), "Failed to reset rules")
|
require.NoError(t, r.Reset(), "Failed to reset rules")
|
||||||
}(r)
|
}(r)
|
||||||
|
|
||||||
@@ -314,16 +313,16 @@ func TestRouter_AddRouteFiltering(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
ruleKey, err := r.AddRouteFiltering(nil, tt.sources, firewall.Network{Prefix: tt.destination}, tt.proto, tt.sPort, tt.dPort, tt.action)
|
ruleKey, err := r.AddFilterRule(nil, tt.sources, firewall.Network{Prefix: tt.destination}, tt.proto, tt.sPort, tt.dPort, tt.action)
|
||||||
require.NoError(t, err, "AddRouteFiltering failed")
|
require.NoError(t, err, "AddFilterRule failed")
|
||||||
|
|
||||||
t.Cleanup(func() {
|
t.Cleanup(func() {
|
||||||
require.NoError(t, r.DeleteRouteRule(ruleKey), "Failed to delete rule")
|
require.NoError(t, r.DeleteFilterRule(ruleKey), "Failed to delete rule")
|
||||||
})
|
})
|
||||||
|
|
||||||
// Check if the rule is in the internal map
|
stored, ok := r.filters[id.RuleID(ruleKey.ID())]
|
||||||
rule, ok := r.rules[ruleKey.ID()]
|
require.True(t, ok, "Rule not found in filters map")
|
||||||
assert.True(t, ok, "Rule not found in internal map")
|
rule := stored.nftRule
|
||||||
|
|
||||||
t.Log("Internal rule expressions:")
|
t.Log("Internal rule expressions:")
|
||||||
for i, expr := range rule.Exprs {
|
for i, expr := range rule.Exprs {
|
||||||
@@ -339,7 +338,7 @@ func TestRouter_AddRouteFiltering(t *testing.T) {
|
|||||||
|
|
||||||
var nftRule *nftables.Rule
|
var nftRule *nftables.Rule
|
||||||
for _, rule := range rules {
|
for _, rule := range rules {
|
||||||
if string(rule.UserData) == ruleKey.ID() {
|
if firewall.RuleID(rule.UserData) == ruleKey.ID() {
|
||||||
nftRule = rule
|
nftRule = rule
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
@@ -367,12 +366,11 @@ func TestNftablesCreateIpSet(t *testing.T) {
|
|||||||
|
|
||||||
defer deleteWorkTable()
|
defer deleteWorkTable()
|
||||||
|
|
||||||
r, err := newRouter(workTable, ifaceMock, iface.DefaultMTU)
|
r := newFamily(workTable, ifaceMock, iface.DefaultMTU)
|
||||||
require.NoError(t, err, "Failed to create router")
|
|
||||||
require.NoError(t, r.init(workTable))
|
require.NoError(t, r.init(workTable))
|
||||||
|
|
||||||
defer func() {
|
defer func() {
|
||||||
require.NoError(t, r.Reset(), "Failed to reset router")
|
require.NoError(t, r.Reset(), "Failed to reset family")
|
||||||
}()
|
}()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
@@ -509,6 +507,58 @@ func TestNftablesCreateIpSet(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestNftablesUpdateSetMergesOverlapping verifies that UpdateSet merges
|
||||||
|
// overlapping prefixes before adding them. An interval set rejects
|
||||||
|
// overlapping elements, so without the merge a batch holding a /32 already
|
||||||
|
// covered by a /24, or a duplicate address as DNS resolution can produce,
|
||||||
|
// would fail.
|
||||||
|
func TestNftablesUpdateSetMergesOverlapping(t *testing.T) {
|
||||||
|
if check() != NFTABLES {
|
||||||
|
t.Skip("nftables not supported on this system")
|
||||||
|
}
|
||||||
|
|
||||||
|
workTable, err := createWorkTable()
|
||||||
|
require.NoError(t, err, "create work table")
|
||||||
|
defer deleteWorkTable()
|
||||||
|
|
||||||
|
r := newFamily(workTable, ifaceMock, iface.DefaultMTU)
|
||||||
|
require.NoError(t, r.init(workTable))
|
||||||
|
defer func() {
|
||||||
|
require.NoError(t, r.Reset(), "reset family")
|
||||||
|
}()
|
||||||
|
|
||||||
|
initial := []netip.Prefix{netip.MustParsePrefix("10.0.0.0/24")}
|
||||||
|
set := firewall.NewPrefixSet(initial)
|
||||||
|
|
||||||
|
created, err := r.createIpSet(set.HashedName(), setInput{prefixes: initial})
|
||||||
|
require.NoError(t, err, "create ip set")
|
||||||
|
require.NotNil(t, created)
|
||||||
|
|
||||||
|
overlapping := []netip.Prefix{
|
||||||
|
netip.MustParsePrefix("192.168.1.0/24"),
|
||||||
|
netip.MustParsePrefix("192.168.1.1/32"),
|
||||||
|
netip.MustParsePrefix("192.168.1.1/32"),
|
||||||
|
}
|
||||||
|
require.NoError(t, r.UpdateSet(set, overlapping), "UpdateSet must merge overlapping prefixes")
|
||||||
|
|
||||||
|
fetchedSet, err := r.conn.GetSetByName(r.workTable, set.HashedName())
|
||||||
|
require.NoError(t, err, "fetch updated set")
|
||||||
|
elements, err := r.conn.GetSetElements(fetchedSet)
|
||||||
|
require.NoError(t, err, "get set elements")
|
||||||
|
|
||||||
|
starts := make(map[string]bool)
|
||||||
|
for _, elem := range elements {
|
||||||
|
if elem.IntervalEnd {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
starts[netip.AddrFrom4(*(*[4]byte)(elem.Key)).String()] = true
|
||||||
|
}
|
||||||
|
// The /32s are covered by the /24, so the update adds one interval and
|
||||||
|
// leaves the one created earlier in place.
|
||||||
|
assert.Equal(t, map[string]bool{"10.0.0.0": true, "192.168.1.0": true}, starts,
|
||||||
|
"merged set must hold the original and the merged interval")
|
||||||
|
}
|
||||||
|
|
||||||
func TestNftablesCreateIpSet_IPv6(t *testing.T) {
|
func TestNftablesCreateIpSet_IPv6(t *testing.T) {
|
||||||
if check() != NFTABLES {
|
if check() != NFTABLES {
|
||||||
t.Skip("nftables not supported on this system")
|
t.Skip("nftables not supported on this system")
|
||||||
@@ -518,11 +568,10 @@ func TestNftablesCreateIpSet_IPv6(t *testing.T) {
|
|||||||
require.NoError(t, err, "Failed to create v6 work table")
|
require.NoError(t, err, "Failed to create v6 work table")
|
||||||
defer deleteWorkTableIPv6()
|
defer deleteWorkTableIPv6()
|
||||||
|
|
||||||
r, err := newRouter(workTable, ifaceMock, iface.DefaultMTU)
|
r := newFamily(workTable, ifaceMock, iface.DefaultMTU)
|
||||||
require.NoError(t, err, "Failed to create router")
|
|
||||||
require.NoError(t, r.init(workTable))
|
require.NoError(t, r.init(workTable))
|
||||||
defer func() {
|
defer func() {
|
||||||
require.NoError(t, r.Reset(), "Failed to reset router")
|
require.NoError(t, r.Reset(), "Failed to reset family")
|
||||||
}()
|
}()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
@@ -748,6 +797,14 @@ func containsPort(exprs []expr.Any, port *firewall.Port, isSource bool) bool {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
case *expr.Lookup:
|
||||||
|
// Multiple discrete ports compile to an anonymous set lookup
|
||||||
|
// rather than a chain of comparisons. The set's id and name are
|
||||||
|
// assigned dynamically, so matching the lookup is enough here;
|
||||||
|
// the set elements are verified separately.
|
||||||
|
if !port.IsRange && len(port.Values) > 1 {
|
||||||
|
portMatchFound = true
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if payloadFound && portMatchFound {
|
if payloadFound && portMatchFound {
|
||||||
return true
|
return true
|
||||||
@@ -861,13 +918,12 @@ func TestRouter_RefreshRulesMap_RemovesStaleEntries(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
defer deleteWorkTable()
|
defer deleteWorkTable()
|
||||||
|
|
||||||
r, err := newRouter(workTable, ifaceMock, iface.DefaultMTU)
|
r := newFamily(workTable, ifaceMock, iface.DefaultMTU)
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, r.init(workTable))
|
require.NoError(t, r.init(workTable))
|
||||||
defer func() { require.NoError(t, r.Reset()) }()
|
defer func() { require.NoError(t, r.Reset()) }()
|
||||||
|
|
||||||
// Add a real rule to the kernel
|
// Add a real rule to the kernel
|
||||||
ruleKey, err := r.AddRouteFiltering(
|
ruleKey, err := r.AddFilterRule(
|
||||||
nil,
|
nil,
|
||||||
[]netip.Prefix{netip.MustParsePrefix("192.168.1.0/24")},
|
[]netip.Prefix{netip.MustParsePrefix("192.168.1.0/24")},
|
||||||
firewall.Network{Prefix: netip.MustParsePrefix("10.0.0.0/24")},
|
firewall.Network{Prefix: netip.MustParsePrefix("10.0.0.0/24")},
|
||||||
@@ -878,11 +934,11 @@ func TestRouter_RefreshRulesMap_RemovesStaleEntries(t *testing.T) {
|
|||||||
)
|
)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
t.Cleanup(func() {
|
t.Cleanup(func() {
|
||||||
require.NoError(t, r.DeleteRouteRule(ruleKey))
|
require.NoError(t, r.DeleteFilterRule(ruleKey))
|
||||||
})
|
})
|
||||||
|
|
||||||
// Inject a stale entry with Handle=0 (simulates store-before-flush failure)
|
// Inject a stale entry with Handle=0 (simulates store-before-flush failure)
|
||||||
staleKey := "stale-rule-that-does-not-exist"
|
staleKey := firewall.RuleID("stale-rule-that-does-not-exist")
|
||||||
r.rules[staleKey] = &nftables.Rule{
|
r.rules[staleKey] = &nftables.Rule{
|
||||||
Table: r.workTable,
|
Table: r.workTable,
|
||||||
Chain: r.chains[chainNameRoutingFw],
|
Chain: r.chains[chainNameRoutingFw],
|
||||||
@@ -902,6 +958,54 @@ func TestRouter_RefreshRulesMap_RemovesStaleEntries(t *testing.T) {
|
|||||||
assert.NotZero(t, realRule.Handle, "real rule should have a valid handle")
|
assert.NotZero(t, realRule.Handle, "real rule should have a valid handle")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestRouter_DeleteRouteRule_RemovesKernelRule verifies a route filter
|
||||||
|
// rule is actually removed from the kernel on delete. The route chain is
|
||||||
|
// not refreshed by Flush, so the stored rule carries a zero handle;
|
||||||
|
// DeleteFilterRule must pull live handles itself before issuing the
|
||||||
|
// delete or the kernel rule leaks. Regression test for that path.
|
||||||
|
func TestRouter_DeleteRouteRule_RemovesKernelRule(t *testing.T) {
|
||||||
|
if check() != NFTABLES {
|
||||||
|
t.Skip("nftables not supported on this system")
|
||||||
|
}
|
||||||
|
|
||||||
|
workTable, err := createWorkTable()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer deleteWorkTable()
|
||||||
|
|
||||||
|
r := newFamily(workTable, ifaceMock, iface.DefaultMTU)
|
||||||
|
require.NoError(t, r.init(workTable))
|
||||||
|
defer func() { require.NoError(t, r.Reset()) }()
|
||||||
|
|
||||||
|
ruleKey, err := r.AddFilterRule(
|
||||||
|
nil,
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("192.168.1.0/24")},
|
||||||
|
firewall.Network{Prefix: netip.MustParsePrefix("10.0.0.0/24")},
|
||||||
|
firewall.ProtocolTCP,
|
||||||
|
nil,
|
||||||
|
&firewall.Port{Values: []uint16{80}},
|
||||||
|
firewall.ActionAccept,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
countKernelRules := func() int {
|
||||||
|
list, err := r.conn.GetRules(r.workTable, r.chains[chainNameRoutingFw])
|
||||||
|
require.NoError(t, err)
|
||||||
|
n := 0
|
||||||
|
for _, rule := range list {
|
||||||
|
if string(rule.UserData) == string(ruleKey.ID()) {
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
require.Equal(t, 1, countKernelRules(), "rule should be present in the kernel after add")
|
||||||
|
|
||||||
|
require.NoError(t, r.DeleteFilterRule(ruleKey))
|
||||||
|
assert.Equal(t, 0, countKernelRules(), "rule must be removed from the kernel after delete")
|
||||||
|
assert.NotContains(t, r.filters, ruleKey.ID(), "filters map entry should be cleared")
|
||||||
|
}
|
||||||
|
|
||||||
func TestRouter_DeleteRouteRule_StaleHandle(t *testing.T) {
|
func TestRouter_DeleteRouteRule_StaleHandle(t *testing.T) {
|
||||||
if check() != NFTABLES {
|
if check() != NFTABLES {
|
||||||
t.Skip("nftables not supported on this system")
|
t.Skip("nftables not supported on this system")
|
||||||
@@ -911,24 +1015,27 @@ func TestRouter_DeleteRouteRule_StaleHandle(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
defer deleteWorkTable()
|
defer deleteWorkTable()
|
||||||
|
|
||||||
r, err := newRouter(workTable, ifaceMock, iface.DefaultMTU)
|
r := newFamily(workTable, ifaceMock, iface.DefaultMTU)
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, r.init(workTable))
|
require.NoError(t, r.init(workTable))
|
||||||
defer func() { require.NoError(t, r.Reset()) }()
|
defer func() { require.NoError(t, r.Reset()) }()
|
||||||
|
|
||||||
// Inject a stale entry with Handle=0
|
// Inject a stale entry with Handle=0
|
||||||
staleKey := "stale-route-rule"
|
staleKey := id.RuleID("stale-route-rule")
|
||||||
r.rules[staleKey] = &nftables.Rule{
|
staleRule := &Rule{
|
||||||
Table: r.workTable,
|
nftRule: &nftables.Rule{
|
||||||
Chain: r.chains[chainNameRoutingFw],
|
Table: r.workTable,
|
||||||
Handle: 0,
|
Chain: r.chains[chainNameRoutingFw],
|
||||||
UserData: []byte(staleKey),
|
Handle: 0,
|
||||||
|
UserData: []byte(staleKey),
|
||||||
|
},
|
||||||
|
id: staleKey,
|
||||||
}
|
}
|
||||||
|
r.filters[staleKey] = staleRule
|
||||||
|
|
||||||
// DeleteRouteRule should not return an error for stale handles
|
// DeleteFilterRule should not return an error for stale handles
|
||||||
err = r.DeleteRouteRule(id.RuleID(staleKey))
|
err = r.DeleteFilterRule(staleRule)
|
||||||
assert.NoError(t, err, "deleting a stale rule should not error")
|
assert.NoError(t, err, "deleting a stale rule should not error")
|
||||||
assert.NotContains(t, r.rules, staleKey, "stale entry should be cleaned up")
|
assert.NotContains(t, r.filters, staleKey, "stale entry should be cleaned up")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRouter_AddNatRule_WithStaleEntry(t *testing.T) {
|
func TestRouter_AddNatRule_WithStaleEntry(t *testing.T) {
|
||||||
@@ -950,7 +1057,7 @@ func TestRouter_AddNatRule_WithStaleEntry(t *testing.T) {
|
|||||||
Masquerade: true,
|
Masquerade: true,
|
||||||
}
|
}
|
||||||
|
|
||||||
rtr := manager.router
|
rtr := manager.family4
|
||||||
|
|
||||||
// First add succeeds
|
// First add succeeds
|
||||||
err = rtr.AddNatRule(pair)
|
err = rtr.AddNatRule(pair)
|
||||||
@@ -960,11 +1067,11 @@ func TestRouter_AddNatRule_WithStaleEntry(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
// Corrupt the handle to simulate stale state
|
// Corrupt the handle to simulate stale state
|
||||||
natRuleKey := firewall.GenKey(firewall.PreroutingFormat, pair)
|
natRuleKey := pair.GenKey(firewall.PreroutingFormat)
|
||||||
if rule, exists := rtr.rules[natRuleKey]; exists {
|
if rule, exists := rtr.rules[natRuleKey]; exists {
|
||||||
rule.Handle = 0
|
rule.Handle = 0
|
||||||
}
|
}
|
||||||
inverseKey := firewall.GenKey(firewall.PreroutingFormat, firewall.GetInversePair(pair))
|
inverseKey := firewall.GetInversePair(pair).GenKey(firewall.PreroutingFormat)
|
||||||
if rule, exists := rtr.rules[inverseKey]; exists {
|
if rule, exists := rtr.rules[inverseKey]; exists {
|
||||||
rule.Handle = 0
|
rule.Handle = 0
|
||||||
}
|
}
|
||||||
@@ -979,7 +1086,7 @@ func TestRouter_AddNatRule_WithStaleEntry(t *testing.T) {
|
|||||||
|
|
||||||
found := 0
|
found := 0
|
||||||
for _, rule := range rules {
|
for _, rule := range rules {
|
||||||
if len(rule.UserData) > 0 && string(rule.UserData) == natRuleKey {
|
if len(rule.UserData) > 0 && firewall.RuleID(rule.UserData) == natRuleKey {
|
||||||
found++
|
found++
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1010,7 +1117,7 @@ func TestCalculateLastIP(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestConvertPrefixesToSet_IPv6(t *testing.T) {
|
func TestConvertPrefixesToSet_IPv6(t *testing.T) {
|
||||||
r := &router{af: afIPv6}
|
r := &family{af: afIPv6}
|
||||||
prefixes := []netip.Prefix{
|
prefixes := []netip.Prefix{
|
||||||
netip.MustParsePrefix("fd00::/64"),
|
netip.MustParsePrefix("fd00::/64"),
|
||||||
netip.MustParsePrefix("2001:db8::1/128"),
|
netip.MustParsePrefix("2001:db8::1/128"),
|
||||||
|
|||||||
@@ -0,0 +1,558 @@
|
|||||||
|
//go:build !android
|
||||||
|
|
||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/google/nftables"
|
||||||
|
"github.com/google/nftables/binaryutil"
|
||||||
|
"github.com/google/nftables/expr"
|
||||||
|
"github.com/hashicorp/go-multierror"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
|
||||||
|
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||||
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
nbnet "github.com/netbirdio/netbird/client/net"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (r *family) AddNatRule(pair firewall.RouterPair) error {
|
||||||
|
if err := r.refreshRulesMap(); err != nil {
|
||||||
|
return fmt.Errorf(refreshRulesMapError, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Resolve every rule's match expressions before queueing any of them: a
|
||||||
|
// message buffered on the shared connection cannot be un-queued, so
|
||||||
|
// returning an error after queueing would leave the next caller's Flush
|
||||||
|
// to commit a rule nothing tracks.
|
||||||
|
var legacyExprs []expr.Any
|
||||||
|
if r.legacyManagement {
|
||||||
|
log.Warnf("This peer is connected to a NetBird Management service with an older version. Allowing all traffic for %s", pair.Destination)
|
||||||
|
|
||||||
|
var err error
|
||||||
|
legacyExprs, err = r.legacyRouteRuleExprs(pair)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("build legacy routing rule: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
inverse := firewall.GetInversePair(pair)
|
||||||
|
var natExprs, inverseExprs []expr.Any
|
||||||
|
if pair.Masquerade {
|
||||||
|
var err error
|
||||||
|
natExprs, err = r.natRuleExprs(pair)
|
||||||
|
if err != nil {
|
||||||
|
r.dropNetworkMatch(legacyExprs)
|
||||||
|
return fmt.Errorf("build nat rule: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
inverseExprs, err = r.natRuleExprs(inverse)
|
||||||
|
if err != nil {
|
||||||
|
r.dropNetworkMatch(legacyExprs)
|
||||||
|
r.dropNetworkMatch(natExprs)
|
||||||
|
return fmt.Errorf("build inverse nat rule: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if legacyExprs != nil {
|
||||||
|
r.queueLegacyRouteRule(pair, legacyExprs)
|
||||||
|
}
|
||||||
|
if pair.Masquerade {
|
||||||
|
r.queueNatRule(pair, natExprs)
|
||||||
|
r.queueNatRule(inverse, inverseExprs)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
r.rollbackRules(pair)
|
||||||
|
return fmt.Errorf("insert rules for %s: %w", pair.Destination, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// rollbackRules cleans up unflushed rules and their set counters after a flush failure.
|
||||||
|
func (r *family) rollbackRules(pair firewall.RouterPair) {
|
||||||
|
keys := []firewall.RuleID{
|
||||||
|
pair.GenKey(firewall.ForwardingFormat),
|
||||||
|
pair.GenKey(firewall.PreroutingFormat),
|
||||||
|
firewall.GetInversePair(pair).GenKey(firewall.PreroutingFormat),
|
||||||
|
}
|
||||||
|
for _, key := range keys {
|
||||||
|
rule, ok := r.rules[key]
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := r.decrementSetCounter(rule); err != nil {
|
||||||
|
log.Warnf("rollback set counter for %s: %v", key, err)
|
||||||
|
}
|
||||||
|
delete(r.rules, key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// natRuleExprs resolves the match expressions of the pair's prerouting
|
||||||
|
// marking rule. It reserves the ipset references the matches need but queues
|
||||||
|
// nothing on the connection, so its error paths leave the connection clean.
|
||||||
|
func (r *family) natRuleExprs(pair firewall.RouterPair) ([]expr.Any, error) {
|
||||||
|
sourceExp, err := r.applyNetwork(pair.Source, nil, true)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("apply source: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
destExp, err := r.applyNetwork(pair.Destination, nil, false)
|
||||||
|
if err != nil {
|
||||||
|
r.dropNetworkMatch(sourceExp)
|
||||||
|
return nil, fmt.Errorf("apply destination: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
op := expr.CmpOpEq
|
||||||
|
if pair.Inverse {
|
||||||
|
op = expr.CmpOpNeq
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs := []expr.Any{
|
||||||
|
&expr.Meta{
|
||||||
|
Key: expr.MetaKeyIIFNAME,
|
||||||
|
Register: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: op,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname(r.wgIface.Name()),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
// We only care about NEW connections to mark them and later identify them in the postrouting chain for masquerading.
|
||||||
|
// Masquerading will take care of the conntrack state, which means we won't need to mark established connections.
|
||||||
|
exprs = append(exprs, getCtNewExprs()...)
|
||||||
|
|
||||||
|
exprs = append(exprs, sourceExp...)
|
||||||
|
exprs = append(exprs, destExp...)
|
||||||
|
|
||||||
|
markValue := nbnet.PreroutingFwmarkMasquerade
|
||||||
|
if pair.Inverse {
|
||||||
|
markValue = nbnet.PreroutingFwmarkMasqueradeReturn
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 1,
|
||||||
|
Data: binaryutil.NativeEndian.PutUint32(markValue),
|
||||||
|
},
|
||||||
|
&expr.Meta{
|
||||||
|
Key: expr.MetaKeyMARK,
|
||||||
|
SourceRegister: true,
|
||||||
|
Register: 1,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
return exprs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// queueNatRule replaces any tracked rule for the pair and queues the new
|
||||||
|
// prerouting marking rule on the connection. Failures are logged rather than
|
||||||
|
// returned: the caller has already queued messages that only a Flush can
|
||||||
|
// commit, so it must not return early.
|
||||||
|
func (r *family) queueNatRule(pair firewall.RouterPair, exprs []expr.Any) {
|
||||||
|
ruleID := pair.GenKey(firewall.PreroutingFormat)
|
||||||
|
|
||||||
|
if _, exists := r.rules[ruleID]; exists {
|
||||||
|
if err := r.removeNatRule(pair); err != nil {
|
||||||
|
// The rule this replaces may still be in the kernel. Keep tracking
|
||||||
|
// it and skip the new one: overwriting the entry would leave the old
|
||||||
|
// rule installed with nothing that can find it again, while keeping
|
||||||
|
// it lets the next update retry the whole replacement.
|
||||||
|
log.Errorf("replace prerouting rule %s: %v", ruleID, err)
|
||||||
|
r.dropNetworkMatch(exprs)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure nat rules come first, so the mark can be overwritten.
|
||||||
|
// Currently overwritten by the dst-type LOCAL rules for redirected traffic.
|
||||||
|
r.rules[ruleID] = r.conn.InsertRule(&nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chains[chainNameManglePrerouting],
|
||||||
|
Exprs: exprs,
|
||||||
|
UserData: []byte(ruleID),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) addPostroutingRules() {
|
||||||
|
// First masquerade rule for traffic coming in from WireGuard interface
|
||||||
|
exprs := []expr.Any{
|
||||||
|
// Match on the first fwmark
|
||||||
|
&expr.Meta{
|
||||||
|
Key: expr.MetaKeyMARK,
|
||||||
|
Register: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkMasquerade),
|
||||||
|
},
|
||||||
|
|
||||||
|
// We need to exclude the loopback interface as this changes the wg proxy port
|
||||||
|
&expr.Meta{
|
||||||
|
Key: expr.MetaKeyOIFNAME,
|
||||||
|
Register: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpNeq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname("lo"),
|
||||||
|
},
|
||||||
|
&expr.Counter{},
|
||||||
|
&expr.Masq{},
|
||||||
|
}
|
||||||
|
|
||||||
|
r.conn.AddRule(&nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chains[chainNameRoutingNat],
|
||||||
|
Exprs: exprs,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Second masquerade rule for traffic going out through WireGuard interface
|
||||||
|
exprs2 := []expr.Any{
|
||||||
|
// Match on the second fwmark
|
||||||
|
&expr.Meta{
|
||||||
|
Key: expr.MetaKeyMARK,
|
||||||
|
Register: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkMasqueradeReturn),
|
||||||
|
},
|
||||||
|
|
||||||
|
// Match WireGuard interface
|
||||||
|
&expr.Meta{
|
||||||
|
Key: expr.MetaKeyOIFNAME,
|
||||||
|
Register: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ifname(r.wgIface.Name()),
|
||||||
|
},
|
||||||
|
&expr.Counter{},
|
||||||
|
&expr.Masq{},
|
||||||
|
}
|
||||||
|
|
||||||
|
r.conn.AddRule(&nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chains[chainNameRoutingNat],
|
||||||
|
Exprs: exprs2,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// addMSSClampingRules adds MSS clamping rules to prevent fragmentation for forwarded traffic.
|
||||||
|
func (r *family) addMSSClampingRules() error {
|
||||||
|
overhead := uint16(ipv4TCPHeaderSize)
|
||||||
|
if r.af.tableFamily == nftables.TableFamilyIPv6 {
|
||||||
|
overhead = ipv6TCPHeaderSize
|
||||||
|
}
|
||||||
|
if r.mtu <= overhead {
|
||||||
|
log.Debugf("MTU %d too small for MSS clamping (overhead %d), skipping", r.mtu, overhead)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
mss := r.mtu - overhead
|
||||||
|
|
||||||
|
exprsOut := []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{unix.IPPROTO_TCP},
|
||||||
|
},
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 13,
|
||||||
|
Len: 1,
|
||||||
|
},
|
||||||
|
&expr.Bitwise{
|
||||||
|
DestRegister: 1,
|
||||||
|
SourceRegister: 1,
|
||||||
|
Len: 1,
|
||||||
|
Mask: []byte{0x02},
|
||||||
|
Xor: []byte{0x00},
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpNeq,
|
||||||
|
Register: 1,
|
||||||
|
Data: []byte{0x00},
|
||||||
|
},
|
||||||
|
&expr.Counter{},
|
||||||
|
&expr.Exthdr{
|
||||||
|
DestRegister: 1,
|
||||||
|
Type: 2,
|
||||||
|
Offset: 2,
|
||||||
|
Len: 2,
|
||||||
|
Op: expr.ExthdrOpTcpopt,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpGt,
|
||||||
|
Register: 1,
|
||||||
|
Data: binaryutil.BigEndian.PutUint16(uint16(mss)),
|
||||||
|
},
|
||||||
|
&expr.Immediate{
|
||||||
|
Register: 1,
|
||||||
|
Data: binaryutil.BigEndian.PutUint16(uint16(mss)),
|
||||||
|
},
|
||||||
|
&expr.Exthdr{
|
||||||
|
SourceRegister: 1,
|
||||||
|
Type: 2,
|
||||||
|
Offset: 2,
|
||||||
|
Len: 2,
|
||||||
|
Op: expr.ExthdrOpTcpopt,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
r.conn.AddRule(&nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chains[chainNameMangleForward],
|
||||||
|
Exprs: exprsOut,
|
||||||
|
})
|
||||||
|
|
||||||
|
return r.conn.Flush()
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildLegacyRouteRuleExpressions(sourceExp, destExp []expr.Any) []expr.Any {
|
||||||
|
exprs := make([]expr.Any, 0, len(sourceExp)+len(destExp)+2)
|
||||||
|
exprs = append(exprs, sourceExp...)
|
||||||
|
exprs = append(exprs, destExp...)
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Counter{},
|
||||||
|
&expr.Verdict{Kind: expr.VerdictAccept},
|
||||||
|
)
|
||||||
|
return exprs
|
||||||
|
}
|
||||||
|
|
||||||
|
// legacyRouteRuleExprs resolves the match expressions of the pair's legacy
|
||||||
|
// forwarding rule, queueing nothing on the connection.
|
||||||
|
func (r *family) legacyRouteRuleExprs(pair firewall.RouterPair) ([]expr.Any, error) {
|
||||||
|
sourceExp, err := r.applyNetwork(pair.Source, nil, true)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("apply source: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
destExp, err := r.applyNetwork(pair.Destination, nil, false)
|
||||||
|
if err != nil {
|
||||||
|
r.dropNetworkMatch(sourceExp)
|
||||||
|
return nil, fmt.Errorf("apply destination: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return buildLegacyRouteRuleExpressions(sourceExp, destExp), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// queueLegacyRouteRule replaces any tracked rule for the pair and queues the
|
||||||
|
// new legacy forwarding rule. Failures are logged for the same reason as in
|
||||||
|
// queueNatRule.
|
||||||
|
func (r *family) queueLegacyRouteRule(pair firewall.RouterPair, exprs []expr.Any) {
|
||||||
|
ruleID := pair.GenKey(firewall.ForwardingFormat)
|
||||||
|
|
||||||
|
if _, exists := r.rules[ruleID]; exists {
|
||||||
|
if err := r.removeLegacyRouteRule(pair); err != nil {
|
||||||
|
// Keep the old rule tracked instead of losing it, as in queueNatRule.
|
||||||
|
log.Errorf("replace legacy forwarding rule %s: %v", ruleID, err)
|
||||||
|
r.dropNetworkMatch(exprs)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
r.rules[ruleID] = r.conn.AddRule(&nftables.Rule{
|
||||||
|
Table: r.workTable,
|
||||||
|
Chain: r.chains[chainNameRoutingFw],
|
||||||
|
Exprs: exprs,
|
||||||
|
UserData: []byte(ruleID),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// removeLegacyRouteRule removes a legacy routing rule for mgmt servers pre route acls
|
||||||
|
func (r *family) removeLegacyRouteRule(pair firewall.RouterPair) error {
|
||||||
|
ruleID := pair.GenKey(firewall.ForwardingFormat)
|
||||||
|
|
||||||
|
rule, exists := r.rules[ruleID]
|
||||||
|
if !exists {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return r.deleteLegacyRuleEntry(ruleID, rule)
|
||||||
|
}
|
||||||
|
|
||||||
|
// deleteLegacyRuleEntry removes one legacy forwarding rule and drops its
|
||||||
|
// ipset references. It also clears stale entries that never got a handle.
|
||||||
|
func (r *family) deleteLegacyRuleEntry(ruleID firewall.RuleID, rule *nftables.Rule) error {
|
||||||
|
if rule.Handle == 0 {
|
||||||
|
log.Warnf("legacy forwarding rule %s has no handle, removing stale entry", ruleID)
|
||||||
|
if err := r.decrementSetCounter(rule); err != nil {
|
||||||
|
log.Warnf("decrement set counter for stale rule %s: %v", ruleID, err)
|
||||||
|
}
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.conn.DelRule(rule); err != nil {
|
||||||
|
return fmt.Errorf("remove legacy forwarding rule %s: %w", ruleID, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
|
||||||
|
if err := r.decrementSetCounter(rule); err != nil {
|
||||||
|
return fmt.Errorf("decrement set counter: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetLegacyManagement returns the route manager's legacy management mode
|
||||||
|
func (r *family) GetLegacyManagement() bool {
|
||||||
|
return r.legacyManagement
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetLegacyManagement sets the route manager to use legacy management mode
|
||||||
|
func (r *family) SetLegacyManagement(isLegacy bool) {
|
||||||
|
r.legacyManagement = isLegacy
|
||||||
|
}
|
||||||
|
|
||||||
|
// RemoveAllLegacyRouteRules removes all legacy routing rules for mgmt servers pre route acls
|
||||||
|
func (r *family) RemoveAllLegacyRouteRules() error {
|
||||||
|
if err := r.refreshRulesMap(); err != nil {
|
||||||
|
return fmt.Errorf(refreshRulesMapError, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var merr *multierror.Error
|
||||||
|
var found bool
|
||||||
|
for k, rule := range r.rules {
|
||||||
|
if !strings.HasPrefix(string(k), firewall.ForwardingFormatPrefix) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
found = true
|
||||||
|
if err := r.deleteLegacyRuleEntry(k, rule); err != nil {
|
||||||
|
merr = multierror.Append(merr, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Commit the queued deletes here instead of leaving them for whichever
|
||||||
|
// caller flushes next: the tracking entries are already gone, so an
|
||||||
|
// uncommitted delete would leave a rule in the kernel that nothing can
|
||||||
|
// find again.
|
||||||
|
if found {
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf(flushError, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) removeNatPreroutingRules() error {
|
||||||
|
table := &nftables.Table{
|
||||||
|
Name: tableNat,
|
||||||
|
Family: r.af.tableFamily,
|
||||||
|
}
|
||||||
|
chain := &nftables.Chain{
|
||||||
|
Name: chainNameNatPrerouting,
|
||||||
|
Table: table,
|
||||||
|
Hooknum: nftables.ChainHookPrerouting,
|
||||||
|
Priority: nftables.ChainPriorityNATDest,
|
||||||
|
Type: nftables.ChainTypeNAT,
|
||||||
|
}
|
||||||
|
rules, err := r.conn.GetRules(table, chain)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("get rules from nat table: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var merr *multierror.Error
|
||||||
|
|
||||||
|
// Delete rules that have our UserData suffix
|
||||||
|
for _, rule := range rules {
|
||||||
|
if len(rule.UserData) == 0 || !strings.HasSuffix(string(rule.UserData), string(dnatSuffix)) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := r.conn.DelRule(rule); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("delete rule %s: %w", rule.UserData, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf(flushError, err))
|
||||||
|
}
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) RemoveNatRule(pair firewall.RouterPair) error {
|
||||||
|
if err := r.refreshRulesMap(); err != nil {
|
||||||
|
return fmt.Errorf(refreshRulesMapError, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var merr *multierror.Error
|
||||||
|
|
||||||
|
if pair.Masquerade {
|
||||||
|
if err := r.removeNatRule(pair); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove prerouting rule: %w", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.removeNatRule(firewall.GetInversePair(pair)); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove inverse prerouting rule: %w", err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.removeLegacyRouteRule(pair); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("remove legacy routing rule: %w", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set counters are decremented in the sub-methods above before flush. If flush fails,
|
||||||
|
// counters will be off until the next successful removal or refresh cycle.
|
||||||
|
if err := r.conn.Flush(); err != nil {
|
||||||
|
merr = multierror.Append(merr, fmt.Errorf("flush remove nat rules %s: %w", pair.Destination, err))
|
||||||
|
}
|
||||||
|
|
||||||
|
return nberrors.FormatErrorOrNil(merr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *family) removeNatRule(pair firewall.RouterPair) error {
|
||||||
|
ruleID := pair.GenKey(firewall.PreroutingFormat)
|
||||||
|
|
||||||
|
rule, exists := r.rules[ruleID]
|
||||||
|
if !exists {
|
||||||
|
log.Debugf("prerouting rule %s not found", ruleID)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if rule.Handle == 0 {
|
||||||
|
log.Warnf("prerouting rule %s has no handle, removing stale entry", ruleID)
|
||||||
|
if err := r.decrementSetCounter(rule); err != nil {
|
||||||
|
log.Warnf("decrement set counter for stale rule %s: %v", ruleID, err)
|
||||||
|
}
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.conn.DelRule(rule); err != nil {
|
||||||
|
return fmt.Errorf("remove prerouting rule %s -> %s: %w", pair.Source, pair.Destination, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugf("removed prerouting rule %s -> %s", pair.Source, pair.Destination)
|
||||||
|
|
||||||
|
delete(r.rules, ruleID)
|
||||||
|
|
||||||
|
if err := r.decrementSetCounter(rule); err != nil {
|
||||||
|
return fmt.Errorf("decrement set counter: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -1,21 +1,26 @@
|
|||||||
package nftables
|
package nftables
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"net"
|
"net/netip"
|
||||||
|
|
||||||
"github.com/google/nftables"
|
"github.com/google/nftables"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Rule to handle management of rules
|
// Rule wraps an installed filter rule (peer or route). Source set
|
||||||
|
// membership is encoded in the rule's expressions; DeleteFilterRule
|
||||||
|
// recovers the set name via findSets so the refcounter can drop the
|
||||||
|
// right reference. mangleRule is set only for peer rules.
|
||||||
type Rule struct {
|
type Rule struct {
|
||||||
nftRule *nftables.Rule
|
nftRule *nftables.Rule
|
||||||
mangleRule *nftables.Rule
|
mangleRule *nftables.Rule
|
||||||
nftSet *nftables.Set
|
// sources is the canonical source list this rule was created for.
|
||||||
ruleID string
|
sources []netip.Prefix
|
||||||
ip net.IP
|
id manager.RuleID
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetRuleID returns the rule id
|
// ID returns the rule id
|
||||||
func (r *Rule) ID() string {
|
func (r *Rule) ID() manager.RuleID {
|
||||||
return r.ruleID
|
return r.id
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,27 @@
|
|||||||
|
//go:build privileged
|
||||||
|
|
||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
)
|
||||||
|
|
||||||
|
func pfx(ip net.IP) []netip.Prefix {
|
||||||
|
if ip == nil {
|
||||||
|
return []netip.Prefix{netip.PrefixFrom(netip.IPv4Unspecified(), 0)}
|
||||||
|
}
|
||||||
|
if ip.IsUnspecified() {
|
||||||
|
if ip.To4() != nil {
|
||||||
|
return []netip.Prefix{netip.PrefixFrom(netip.IPv4Unspecified(), 0)}
|
||||||
|
}
|
||||||
|
return []netip.Prefix{netip.PrefixFrom(netip.IPv6Unspecified(), 0)}
|
||||||
|
}
|
||||||
|
a, ok := netip.AddrFromSlice(ip)
|
||||||
|
if !ok {
|
||||||
|
panic(fmt.Sprintf("invalid IP length: %d", len(ip)))
|
||||||
|
}
|
||||||
|
a = a.Unmap()
|
||||||
|
return []netip.Prefix{netip.PrefixFrom(a, a.BitLen())}
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user